diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..d7989c7 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,11 @@ +# Changelog + +## 0.2.1 + +- 适配 MCP Python SDK 2.x:低层 `Server` 改为构造参数 `on_list_tools` / `on_call_tool`,返回 `ListToolsResult` / `CallToolResult`,工具列表变更通知改走 `ctx.session.send_tool_list_changed()`。 +- 依赖改为 `mcp>=2.2,<3`,避免再解析到已删除 `Server.list_tools` 的版本,也避免下一个大版本无上限装崩。 +- `__version__` 与包版本对齐为 0.2.1。 + +## 0.2.0 + +- 增加 Link 模式认证(`--link`):远程客户端通过授权链接与授权码完成登录,不依赖本机浏览器回调。 diff --git a/claude.md b/claude.md index 294fc87..5304e24 100644 --- a/claude.md +++ b/claude.md @@ -1,5 +1,7 @@ # Uno MCP Stdio - Claude 项目指南 +当前版本 **0.2.1**。运行时依赖 `mcp>=2.2,<3`。低层 Server 用 `on_list_tools` / `on_call_tool`(mcp 2.x 已删除装饰器 `list_tools`)。仓库没有发布 workflow,合并不会自动发到 PyPI。 + ## 项目概述 `uno-mcp-stdio` 是 Uno MCP Gateway 的本地 stdio 代理客户端。它解决了不支持 OAuth 认证的 MCP 客户端(如 Manus、Cherry Studio)无法连接需要认证的 MCP 服务器的问题。 @@ -57,19 +59,12 @@ uno-mcp-stdio/ ### 1. stdio_server.py - MCP Server 实现 -使用 MCP Python SDK 实现 stdio 传输的 server: - -```python -class UnoStdioServer: - # 处理 tools/list - 代理到 gateway 获取工具列表 - # 处理 tools/call - 代理到 gateway 执行工具 - # 处理 uno_auth_required - 启动 OAuth 认证流程 -``` +使用 MCP Python SDK 2.x 低层 `Server`(`on_list_tools` / `on_call_tool`)实现 stdio server,代理到远程 gateway。本地工具名是 `uno_auth`(登录、退出、状态;Link 模式带 `code`)。 关键点: -- 如果未认证,`tools/list` 返回一个 `uno_auth_required` 工具 -- 用户调用该工具触发 OAuth 认证流程 -- 认证成功后,重新调用 `tools/list` 获取真实工具列表 +- `tools/list` 始终在列表开头放 `uno_auth`,其余工具来自 gateway +- 未登录时 `tools/call` 除 `uno_auth` 外返回需要认证 +- 认证成功后通过 `ctx.session.send_tool_list_changed()` 通知客户端刷新 ### 2. token_manager.py - Token 管理 diff --git a/pyproject.toml b/pyproject.toml index 4fc7be2..cbd8e12 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "uno-mcp-stdio" -version = "0.1.9" +version = "0.2.1" description = "Uno MCP Stdio Client - Local stdio proxy for Uno MCP Gateway with OAuth authentication" readme = "README.md" requires-python = ">=3.11" @@ -21,7 +21,7 @@ classifiers = [ ] dependencies = [ - "mcp>=1.0.0", + "mcp>=2.2,<3", "httpx>=0.27.0", "pydantic>=2.5.0", "pydantic-settings>=2.1.0", diff --git a/src/uno_mcp_stdio/__init__.py b/src/uno_mcp_stdio/__init__.py index 651990a..c087282 100644 --- a/src/uno_mcp_stdio/__init__.py +++ b/src/uno_mcp_stdio/__init__.py @@ -4,5 +4,5 @@ 为不支持 OAuth 认证的 MCP 客户端提供本地代理。 """ -__version__ = "0.1.3" +__version__ = "0.2.1" diff --git a/src/uno_mcp_stdio/auth/token_manager.py b/src/uno_mcp_stdio/auth/token_manager.py index 577325b..e23897f 100644 --- a/src/uno_mcp_stdio/auth/token_manager.py +++ b/src/uno_mcp_stdio/auth/token_manager.py @@ -92,6 +92,45 @@ def from_dict(cls, data: Dict[str, Any]) -> "ClientRegistration": ) +@dataclass +class PendingAuthSession: + """ + 待完成的认证会话(用于 Link 模式) + + Link 模式下,用户需要分两步完成认证: + 1. 获取认证链接 + 2. 输入授权码 + + 这个类存储第一步生成的 PKCE 参数,供第二步使用。 + """ + code_verifier: str + code_challenge: str + state: str + redirect_uri: str + client_id: str + created_at: float # Unix timestamp + + def is_expired(self, timeout_seconds: int = 600) -> bool: + """检查会话是否过期(默认 10 分钟)""" + return time.time() > (self.created_at + timeout_seconds) + + def to_dict(self) -> Dict[str, Any]: + """转换为字典""" + return asdict(self) + + @classmethod + def from_dict(cls, data: Dict[str, Any]) -> "PendingAuthSession": + """从字典创建""" + return cls( + code_verifier=data["code_verifier"], + code_challenge=data["code_challenge"], + state=data["state"], + redirect_uri=data["redirect_uri"], + client_id=data["client_id"], + created_at=data["created_at"] + ) + + class TokenManager: """Token 管理器""" @@ -101,6 +140,9 @@ def __init__(self): self._oauth_metadata: Optional[OAuthMetadata] = None self._client_registration: Optional[ClientRegistration] = None self._client_registration_path = self._credentials_path.parent / "client.json" + # Link 模式相关 + self._pending_session: Optional[PendingAuthSession] = None + self._pending_session_path = self._credentials_path.parent / "pending_session.json" def _log(self, message: str): """输出日志到 stderr(避免干扰 stdio 通信)""" @@ -251,6 +293,64 @@ async def ensure_client_registered(self, redirect_uri: str) -> Optional[str]: self._log("动态注册失败,使用默认 client_id") return settings.oauth_client_id + # ==================== Pending Session Management (Link Mode) ==================== + + def save_pending_session(self, session: PendingAuthSession): + """ + 保存待完成的认证会话(Link 模式) + + 在用户获取认证链接后,保存 PKCE 参数,等待用户输入授权码。 + """ + self._pending_session = session + + self._pending_session_path.parent.mkdir(parents=True, exist_ok=True) + + with open(self._pending_session_path, "w") as f: + json.dump(session.to_dict(), f, indent=2) + + self._log(f"已保存 pending session: state={session.state[:8]}...") + + def load_pending_session(self) -> Optional[PendingAuthSession]: + """加载待完成的认证会话""" + if self._pending_session: + if not self._pending_session.is_expired(): + return self._pending_session + else: + self._log("内存中的 pending session 已过期") + self._pending_session = None + + if not self._pending_session_path.exists(): + return None + + try: + with open(self._pending_session_path, "r") as f: + data = json.load(f) + session = PendingAuthSession.from_dict(data) + + if session.is_expired(): + self._log("文件中的 pending session 已过期,清除") + self.clear_pending_session() + return None + + self._pending_session = session + self._log(f"已加载 pending session: state={session.state[:8]}...") + return session + except Exception as e: + self._log(f"加载 pending session 失败: {e}") + return None + + def clear_pending_session(self): + """清除待完成的认证会话""" + self._pending_session = None + if self._pending_session_path.exists(): + self._pending_session_path.unlink() + self._log("已清除 pending session") + + def has_pending_session(self) -> bool: + """检查是否有待完成的认证会话""" + session = self.load_pending_session() + return session is not None + # ==================== Credentials Management ==================== def load_credentials(self) -> Optional[Credentials]: @@ -466,6 +566,92 @@ def open_auth_url(self, url: str) -> bool: except Exception as e: self._log(f"无法打开浏览器: {e}") return False + + # ==================== Link Mode Methods ==================== + + async def create_link_mode_session(self) -> Optional[Dict[str, str]]: + """ + 创建 Link 模式认证会话 + + 生成认证 URL 并保存 PKCE 参数,返回认证信息供用户使用。 + + Returns: + { + "auth_url": "认证链接", + "state": "会话标识(可选,用于验证)" + } + """ + # 生成 PKCE 参数 + code_verifier, code_challenge = self.generate_pkce() + state = self.generate_state() + + # 使用固定的回调 URL(MCPMarket 提供的授权码显示页面) + redirect_uri = settings.link_mode_callback_url + + # 确保客户端已注册 + client_id = await self.ensure_client_registered(redirect_uri) + if not client_id: + self._log("客户端注册失败") + return None + + # 构建认证 URL + auth_url = await self.build_auth_url(redirect_uri, state, code_challenge, client_id) + if not auth_url: + self._log("构建认证 URL 失败") + return None + + # 保存 pending session + session = PendingAuthSession( + code_verifier=code_verifier, + code_challenge=code_challenge, + state=state, + redirect_uri=redirect_uri, + client_id=client_id, + created_at=time.time() + ) + self.save_pending_session(session) + + self._log(f"Link 模式会话已创建: auth_url={auth_url[:50]}...") + + return { + "auth_url": auth_url, + "state": state + } + + async def complete_link_mode_auth(self, code: str) -> Optional[Credentials]: + """ + 完成 Link 模式认证 + + 使用用户提供的授权码交换 token。 + + Args: + code: 用户从认证页面获取的授权码 + + Returns: + 认证成功返回 Credentials,失败返回 None + """ + # 加载 pending session + session = self.load_pending_session() + if not session: + self._log("没有待完成的认证会话") + return None + + # 交换 token + credentials = await self.exchange_code_for_token( + code=code, + code_verifier=session.code_verifier, + redirect_uri=session.redirect_uri, + client_id=session.client_id + ) + + if credentials: + # 清除 pending session + self.clear_pending_session() + self._log("Link 模式认证成功") + return credentials + else: + self._log("Link 模式认证失败:token 交换失败") + return None # 全局实例 diff --git a/src/uno_mcp_stdio/config.py b/src/uno_mcp_stdio/config.py index 2d5427b..1783733 100644 --- a/src/uno_mcp_stdio/config.py +++ b/src/uno_mcp_stdio/config.py @@ -70,6 +70,18 @@ class Settings(BaseSettings): description="等待 OAuth 回调超时时间(秒)" ) + # Link 模式配置(用于远程服务器场景,如 Manus) + link_mode_callback_url: str = Field( + default="https://mcpmarket.cn/oauth/code-display", + description="Link 模式下的回调 URL,该页面会显示授权码供用户复制" + ) + + # 认证模式 + auth_mode: str = Field( + default="auto", + description="认证模式: auto(自动检测), local(本地模式), link(链接模式)" + ) + def get_credentials_path(self) -> Path: """获取 credentials 文件的完整路径""" path = Path(self.credentials_path).expanduser() diff --git a/src/uno_mcp_stdio/main.py b/src/uno_mcp_stdio/main.py index d70fda8..ec8f99c 100644 --- a/src/uno_mcp_stdio/main.py +++ b/src/uno_mcp_stdio/main.py @@ -2,40 +2,98 @@ Uno MCP Stdio - 入口文件 提供命令行入口,启动 stdio server。 + +用法: + uvx uno-mcp-stdio # 本地模式(默认) + uvx uno-mcp-stdio --link # 链接模式(用于 Manus 等远程服务器场景) """ import sys +import argparse import asyncio from .config import settings -def print_banner(): +def print_banner(link_mode: bool = False): """打印启动横幅(到 stderr,避免干扰 stdio)""" - banner = """ + mode_str = "Link Mode (远程)" if link_mode else "Local Mode (本地)" + banner = f""" ╔═══════════════════════════════════════════════════════════╗ ║ Uno MCP Stdio ║ ║ Local proxy for Uno MCP Gateway ║ +║ [{mode_str}] ╚═══════════════════════════════════════════════════════════╝ """ print(banner, file=sys.stderr) print(f" Gateway: {settings.gateway_url}", file=sys.stderr) print(f" Credentials: {settings.get_credentials_path()}", file=sys.stderr) + print(f" Auth Mode: {'link' if link_mode else 'local'}", file=sys.stderr) print(f" Debug: {settings.debug}", file=sys.stderr) print("", file=sys.stderr) +def parse_args(): + """解析命令行参数""" + parser = argparse.ArgumentParser( + description="Uno MCP Stdio - Local proxy for Uno MCP Gateway", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +认证模式说明: + 本地模式 (默认): 会弹出浏览器完成认证,适用于本地运行 + 链接模式 (--link): 返回认证链接,用户在其他设备完成认证后输入授权码 + 适用于 Manus 等远程服务器场景 + +示例配置 (Manus/Cherry Studio): + { + "mcpServers": { + "uno": { + "command": "uvx", + "args": ["uno-mcp-stdio", "--link"] + } + } + } +""" + ) + + parser.add_argument( + "--link", "-l", + action="store_true", + help="启用链接模式(用于远程服务器,如 Manus)。认证时返回链接,用户手动完成后输入授权码" + ) + + parser.add_argument( + "--debug", "-d", + action="store_true", + help="启用调试模式,输出详细日志" + ) + + return parser.parse_args() + + def main(): """主入口函数""" + # 解析命令行参数 + args = parse_args() + + # 设置调试模式 + if args.debug: + settings.debug = True + + # 设置认证模式 + link_mode = args.link + if link_mode: + settings.auth_mode = "link" + # 打印横幅 if settings.debug: - print_banner() + print_banner(link_mode) - # 导入并运行服务器 + # 导入并运行服务器(传入 link_mode) from .stdio_server import run_server try: - asyncio.run(run_server()) + asyncio.run(run_server(link_mode=link_mode)) except KeyboardInterrupt: print("\n[UnoStdio] 收到中断信号,正在退出...", file=sys.stderr) sys.exit(0) diff --git a/src/uno_mcp_stdio/stdio_server.py b/src/uno_mcp_stdio/stdio_server.py index 45772d6..edc6166 100644 --- a/src/uno_mcp_stdio/stdio_server.py +++ b/src/uno_mcp_stdio/stdio_server.py @@ -10,15 +10,19 @@ import asyncio from typing import Optional -from mcp.server import Server, NotificationOptions +from mcp.server import Server, NotificationOptions, ServerRequestContext from mcp.server.stdio import stdio_server from mcp.types import ( - Tool, - TextContent, + CallToolRequestParams, CallToolResult, ListToolsResult, + PaginatedRequestParams, + TextContent, + Tool, ) +from . import __version__ + from .config import settings from .auth import token_manager, CallbackServer from .gateway import gateway_proxy, AuthenticationRequired, GatewayError @@ -27,11 +31,23 @@ class UnoStdioServer: """Uno MCP Stdio Server""" - def __init__(self): - self.server = Server("uno-mcp-stdio") + def __init__(self, link_mode: bool = False): + """ + 初始化 Uno MCP Stdio Server + + Args: + link_mode: 是否使用链接模式(用于 Manus 等远程服务器场景) + """ self._authenticated = False self._tools_cache: Optional[list] = None - self._setup_handlers() + self._link_mode = link_mode + # mcp 2.x:处理器在构造时注册,不再有 list_tools/call_tool 装饰器。 + self.server = Server( + "uno-mcp-stdio", + version=__version__, + on_list_tools=self._on_list_tools, + on_call_tool=self._on_call_tool, + ) def _get_notification_options(self) -> NotificationOptions: """获取通知选项,声明支持 tools_changed 通知""" @@ -45,159 +61,189 @@ def _log(self, message: str): """输出日志到 stderr(避免干扰 stdio 通信)""" print(f"[UnoStdio] {message}", file=sys.stderr, flush=True) - def _setup_handlers(self): - """设置 MCP 请求处理器""" - - @self.server.list_tools() - async def handle_list_tools() -> list[Tool]: - """处理 tools/list 请求""" - self._log("收到 tools/list 请求") - - # 检查是否已认证 - is_authenticated = self._authenticated or bool(await token_manager.get_valid_token()) - - # 获取工具列表(proxy 内部会自动处理 default token) - try: - response = await gateway_proxy.list_tools() - - if "result" in response and "tools" in response["result"]: - tools_data = response["result"]["tools"] - self._tools_cache = tools_data - - # 转换为 MCP Tool 对象 - tools = [] - - # 始终在列表开头添加认证工具(用于首次认证、重新认证、退出登录等) - if is_authenticated: - auth_description = "🔐 认证管理工具。当前状态:✅ 已登录。支持的操作:login(重新登录)、logout(退出登录)、status(查看状态)" - else: - auth_description = "🔐 认证管理工具。当前状态:❌ 未登录。请调用此工具完成认证后才能使用其他工具。支持的操作:login(登录)、status(查看状态)" - + def _auth_input_schema(self) -> dict: + """uno_auth 的入参。Link 模式多一个 code。""" + properties: dict = { + "action": { + "type": "string", + "enum": ["login", "logout", "status"], + "description": "操作类型:login(登录/重新登录)、logout(退出登录)、status(查看状态)。默认为 login", + } + } + if self._link_mode: + properties["code"] = { + "type": "string", + "description": "授权码。在链接模式下,用户访问认证链接完成授权后,将页面显示的授权码填入此参数", + } + return {"type": "object", "properties": properties, "required": []} + + @staticmethod + def _tool_schema(raw: object) -> dict: + """协议要求 input_schema 根上有 type=object。""" + if not isinstance(raw, dict): + return {"type": "object"} + if "type" not in raw: + return {"type": "object", **raw} + return raw + + async def _notify_tools_changed(self, session) -> None: + """认证状态变化后通知客户端刷新工具列表。""" + if session is None: + return + try: + await session.send_tool_list_changed() + self._log("已发送 tools/list_changed 通知") + except Exception as e: + self._log(f"发送通知失败(客户端可能不支持): {e}") + + async def _on_list_tools( + self, + ctx: ServerRequestContext, + params: PaginatedRequestParams | None, + ) -> ListToolsResult: + """处理 tools/list。返回完整 ListToolsResult,不再依赖 SDK 自动包装。""" + del ctx, params + self._log("收到 tools/list 请求") + + is_authenticated = self._authenticated or bool(await token_manager.get_valid_token()) + + try: + response = await gateway_proxy.list_tools() + + if "result" in response and "tools" in response["result"]: + tools_data = response["result"]["tools"] + self._tools_cache = tools_data + tools = [] + has_pending = token_manager.has_pending_session() + + if is_authenticated: + auth_description = "🔐 认证管理工具。当前状态:✅ 已登录。支持的操作:login(重新登录)、logout(退出登录)、status(查看状态)" + elif has_pending: + auth_description = "🔐 认证管理工具。当前状态:⏳ 等待输入授权码。请将认证页面显示的授权码通过 code 参数传入完成认证" + else: + auth_description = "🔐 认证管理工具。当前状态:❌ 未登录。请调用此工具完成认证后才能使用其他工具。支持的操作:login(登录)、status(查看状态)" + + tools.append(Tool( + name="uno_auth", + description=auth_description, + input_schema=self._auth_input_schema(), + )) + + for t in tools_data: tools.append(Tool( - name="uno_auth", - description=auth_description, - inputSchema={ - "type": "object", - "properties": { - "action": { - "type": "string", - "enum": ["login", "logout", "status"], - "description": "操作类型:login(登录/重新登录)、logout(退出登录)、status(查看状态)。默认为 login" - } - }, - "required": [] - } + name=t["name"], + description=t.get("description", ""), + input_schema=self._tool_schema( + t.get("inputSchema", t.get("input_schema")) + ), )) - - for t in tools_data: - tools.append(Tool( - name=t["name"], - description=t.get("description", ""), - inputSchema=t.get("inputSchema", {"type": "object"}) - )) - - self._log(f"返回 {len(tools)} 个工具 (已认证: {is_authenticated})") - return tools - else: - self._log(f"Gateway 返回格式异常: {response}") - return [] - - except GatewayError as e: - self._log(f"获取工具列表失败: {e}") - # 如果获取失败,返回认证工具 - return [ - Tool( - name="uno_auth", - description="🔐 认证管理工具。请调用此工具获取认证链接。支持的操作:login(登录)、logout(退出)、status(查看状态)", - inputSchema={ - "type": "object", - "properties": { - "action": { - "type": "string", - "enum": ["login", "logout", "status"], - "description": "操作类型:login(登录/重新登录)、logout(退出登录)、status(查看状态)。默认为 login" - } - }, - "required": [] - } - ) - ] - - @self.server.call_tool() - async def handle_call_tool(name: str, arguments: dict) -> list[TextContent]: - """处理 tools/call 请求""" - self._log(f"收到 tools/call 请求: {name}") - - # 处理认证请求 - if name == "uno_auth": - action = arguments.get("action", "login") - return await self._handle_auth_request(action=action) - - # 检查认证 - try: - await self._ensure_authenticated() - except AuthenticationRequired: - return [TextContent( - type="text", - text=json.dumps({ - "error": "authentication_required", - "message": "需要认证,请先调用 uno_auth 工具" - }, ensure_ascii=False) - )] - - # 代理到 gateway - try: - response = await gateway_proxy.call_tool(name, arguments, request_id=1) - - if "result" in response: - result = response["result"] - # 返回工具调用结果 - if "content" in result: - contents = [] - for item in result["content"]: - if item.get("type") == "text": - contents.append(TextContent( - type="text", - text=item.get("text", "") - )) - return contents - else: - return [TextContent( - type="text", - text=json.dumps(result, ensure_ascii=False, indent=2) - )] - elif "error" in response: - return [TextContent( - type="text", - text=json.dumps({ - "error": response["error"].get("code"), - "message": response["error"].get("message") - }, ensure_ascii=False) - )] - else: - return [TextContent( - type="text", - text=json.dumps(response, ensure_ascii=False) - )] - - except AuthenticationRequired: - token_manager.clear_credentials() - self._authenticated = False - return [TextContent( - type="text", - text=json.dumps({ - "error": "authentication_expired", - "message": "认证已过期,请重新调用 uno_auth 工具" - }, ensure_ascii=False) - )] - except GatewayError as e: - return [TextContent( + + self._log(f"返回 {len(tools)} 个工具 (已认证: {is_authenticated})") + return ListToolsResult(tools=tools) + + self._log(f"Gateway 返回格式异常: {response}") + return ListToolsResult(tools=[]) + + except GatewayError as e: + self._log(f"获取工具列表失败: {e}") + return ListToolsResult(tools=[ + Tool( + name="uno_auth", + description="🔐 认证管理工具。请调用此工具获取认证链接。支持的操作:login(登录)、logout(退出)、status(查看状态)", + input_schema=self._auth_input_schema(), + ) + ]) + + async def _on_call_tool( + self, + ctx: ServerRequestContext, + params: CallToolRequestParams, + ) -> CallToolResult: + """处理 tools/call。异常转成 is_error 结果,避免变成 JSON-RPC 协议错误。""" + try: + return await self._dispatch_call(ctx, params) + except Exception as e: + self._log(f"tools/call 未处理异常: {e}") + return CallToolResult( + content=[TextContent( type="text", text=json.dumps({ - "error": "gateway_error", - "message": str(e) - }, ensure_ascii=False) - )] + "error": "internal_error", + "message": str(e), + }, ensure_ascii=False), + )], + is_error=True, + ) + + async def _dispatch_call( + self, + ctx: ServerRequestContext, + params: CallToolRequestParams, + ) -> CallToolResult: + name = params.name + arguments = params.arguments or {} + self._log(f"收到 tools/call 请求: {name}") + + if name == "uno_auth": + action = arguments.get("action", "login") + code = arguments.get("code") + contents = await self._handle_auth_request( + action=action, code=code, session=ctx.session + ) + return CallToolResult(content=contents, is_error=False) + + try: + await self._ensure_authenticated() + except AuthenticationRequired: + return self._text_result({ + "error": "authentication_required", + "message": "需要认证,请先调用 uno_auth 工具", + }) + + try: + response = await gateway_proxy.call_tool(name, arguments, request_id=1) + + if "result" in response: + result = response["result"] + if "content" in result: + contents = [] + for item in result["content"]: + if item.get("type") == "text": + contents.append(TextContent( + type="text", + text=item.get("text", ""), + )) + return CallToolResult(content=contents, is_error=False) + return self._text_result(result, indent=2) + if "error" in response: + return self._text_result({ + "error": response["error"].get("code"), + "message": response["error"].get("message"), + }) + return self._text_result(response) + + except AuthenticationRequired: + token_manager.clear_credentials() + self._authenticated = False + return self._text_result({ + "error": "authentication_expired", + "message": "认证已过期,请重新调用 uno_auth 工具", + }) + except GatewayError as e: + return self._text_result({ + "error": "gateway_error", + "message": str(e), + }) + + @staticmethod + def _text_result(payload: object, indent: int | None = None) -> CallToolResult: + return CallToolResult( + content=[TextContent( + type="text", + text=json.dumps(payload, ensure_ascii=False, indent=indent), + )], + is_error=False, + ) async def _ensure_authenticated(self): """确保已认证""" @@ -212,7 +258,12 @@ async def _ensure_authenticated(self): raise AuthenticationRequired("需要认证") - async def _handle_auth_request(self, action: str = "login") -> list[TextContent]: + async def _handle_auth_request( + self, + action: str = "login", + code: str = None, + session=None, + ) -> list[TextContent]: """ 处理认证请求 @@ -221,12 +272,15 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] - login: 登录或重新登录 - logout: 退出登录 - status: 查看认证状态 + code: 授权码(Link 模式下使用) """ - self._log(f"处理认证请求: action={action}") + self._log(f"处理认证请求: action={action}, code={'***' if code else 'None'}, link_mode={self._link_mode}") # 处理状态查询 if action == "status": token = await token_manager.get_valid_token() + has_pending = token_manager.has_pending_session() + if token: return [TextContent( type="text", @@ -236,6 +290,15 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] "hint": "可使用 action='logout' 退出登录,或 action='login' 重新登录" }, ensure_ascii=False) )] + elif has_pending: + return [TextContent( + type="text", + text=json.dumps({ + "status": "pending", + "message": "⏳ 等待输入授权码", + "hint": "请访问认证链接完成授权,然后将页面显示的授权码通过 code 参数传入" + }, ensure_ascii=False) + )] else: return [TextContent( type="text", @@ -251,17 +314,12 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] token = await token_manager.get_valid_token() if token: token_manager.clear_credentials() + token_manager.clear_pending_session() # 同时清除 pending session self._authenticated = False self._log("用户已退出登录") - # 发送工具列表变更通知 - try: - session = self.server.request_context.session - await session.send_tool_list_changed() - self._log("已发送 tools/list_changed 通知") - except Exception as e: - self._log(f"发送通知失败(客户端可能不支持): {e}") - + await self._notify_tools_changed(session) + return [TextContent( type="text", text=json.dumps({ @@ -280,6 +338,123 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] )] # 处理登录请求 (action == "login" 或其他) + + # Link 模式:如果提供了 code,尝试完成认证 + if self._link_mode and code: + return await self._handle_link_mode_complete(code, session=session) + + # Link 模式:没有 code,生成认证链接 + if self._link_mode: + return await self._handle_link_mode_start() + + # 本地模式:原有的浏览器认证流程 + return await self._handle_local_mode_auth(session=session) + + async def _handle_link_mode_start(self) -> list[TextContent]: + """ + Link 模式:生成认证链接 + + 返回认证 URL,用户需要在自己的设备上访问完成认证, + 然后将页面显示的授权码传回。 + """ + self._log("Link 模式:生成认证链接") + + # 如果已有 token,先清除(实现重新登录) + existing_token = await token_manager.get_valid_token() + if existing_token: + self._log("检测到已有 token,清除后重新认证") + token_manager.clear_credentials() + self._authenticated = False + + # 清除可能存在的旧 pending session + token_manager.clear_pending_session() + + # 创建新的认证会话 + session_info = await token_manager.create_link_mode_session() + + if not session_info: + return [TextContent( + type="text", + text=json.dumps({ + "status": "failed", + "error": "session_creation_failed", + "message": "❌ 创建认证会话失败,请检查网络连接" + }, ensure_ascii=False) + )] + + auth_url = session_info["auth_url"] + self._log(f"认证链接已生成: {auth_url[:50]}...") + + return [TextContent( + type="text", + text=json.dumps({ + "status": "link_generated", + "message": "🔗 请复制以下链接到浏览器完成认证", + "auth_url": auth_url, + "instructions": [ + "1. 复制上面的 auth_url 链接", + "2. 在浏览器中打开该链接", + "3. 在 MCPMarket 完成登录/授权", + "4. 授权完成后,页面会显示一个授权码", + "5. 将授权码复制,再次调用此工具并设置 code 参数", + " 例如:uno_auth(code='你的授权码')" + ], + "next_step": "获取授权码后,调用 uno_auth(code='授权码') 完成认证" + }, ensure_ascii=False, indent=2) + )] + + async def _handle_link_mode_complete(self, code: str, session=None) -> list[TextContent]: + """ + Link 模式:使用授权码完成认证 + + Args: + code: 用户从认证页面获取的授权码 + """ + self._log(f"Link 模式:使用授权码完成认证") + + # 检查是否有 pending session + if not token_manager.has_pending_session(): + return [TextContent( + type="text", + text=json.dumps({ + "status": "failed", + "error": "no_pending_session", + "message": "❌ 没有待完成的认证会话,请先调用 uno_auth() 获取认证链接" + }, ensure_ascii=False) + )] + + # 使用授权码完成认证 + credentials = await token_manager.complete_link_mode_auth(code) + + if credentials: + self._authenticated = True + self._log("Link 模式认证成功!") + + await self._notify_tools_changed(session) + + return [TextContent( + type="text", + text=json.dumps({ + "status": "success", + "message": "✅ 认证成功!现在可以使用 Uno 的工具了。" + }, ensure_ascii=False) + )] + else: + return [TextContent( + type="text", + text=json.dumps({ + "status": "failed", + "error": "token_exchange_failed", + "message": "❌ 授权码验证失败,请检查授权码是否正确,或重新获取认证链接" + }, ensure_ascii=False) + )] + + async def _handle_local_mode_auth(self, session=None) -> list[TextContent]: + """ + 本地模式:使用浏览器完成认证(原有流程) + """ + self._log("本地模式:启动浏览器认证流程") + # 如果已有 token,先清除(实现重新登录) existing_token = await token_manager.get_valid_token() if existing_token: @@ -329,27 +504,9 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] # 尝试自动打开浏览器 browser_opened = token_manager.open_auth_url(auth_url) - # 返回认证信息 - auth_message = { - "status": "authentication_required", - "message": "请在浏览器中完成认证", - "auth_url": auth_url, - "browser_opened": browser_opened, - "instructions": [ - "1. 点击上面的链接或复制到浏览器打开", - "2. 在 MCPMarket 完成登录/授权", - "3. 授权后页面会自动关闭", - "4. 返回这里继续使用" - ] - } - self._log(f"认证 URL: {auth_url}") self._log("等待用户完成认证...") - # 先返回认证链接 - # 注意:这里需要异步等待回调,但 MCP 的 call_tool 是同步返回的 - # 所以我们需要在后台等待,同时返回信息给用户 - # 等待回调 callback_data = callback_server.wait_for_callback() callback_server.stop() @@ -375,9 +532,9 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] )] # 交换 token - code = callback_data.get("code") + auth_code = callback_data.get("code") credentials = await token_manager.exchange_code_for_token( - code=code, + code=auth_code, code_verifier=code_verifier, redirect_uri=redirect_uri, client_id=client_id @@ -387,14 +544,8 @@ async def _handle_auth_request(self, action: str = "login") -> list[TextContent] self._authenticated = True self._log("认证成功!") - # 发送工具列表变更通知,让客户端刷新工具列表 - try: - session = self.server.request_context.session - await session.send_tool_list_changed() - self._log("已发送 tools/list_changed 通知") - except Exception as e: - self._log(f"发送通知失败(客户端可能不支持): {e}") - + await self._notify_tools_changed(session) + return [TextContent( type="text", text=json.dumps({ @@ -441,8 +592,13 @@ async def run(self): self._log("Uno MCP Stdio Server 已关闭") -async def run_server(): - """运行服务器入口""" - server = UnoStdioServer() +async def run_server(link_mode: bool = False): + """ + 运行服务器入口 + + Args: + link_mode: 是否使用链接模式(用于 Manus 等远程服务器场景) + """ + server = UnoStdioServer(link_mode=link_mode) await server.run() diff --git a/tests/test_stdio_handlers.py b/tests/test_stdio_handlers.py new file mode 100644 index 0000000..bb292d7 --- /dev/null +++ b/tests/test_stdio_handlers.py @@ -0,0 +1,128 @@ +"""mcp 2.x 低层 handler 的返回形状。不访问网络,也不写真实凭据。""" + +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock + +from mcp.types import CallToolRequestParams + +from uno_mcp_stdio import stdio_server as mod +from uno_mcp_stdio.gateway import GatewayError + + +def _ctx(): + session = MagicMock() + session.send_tool_list_changed = AsyncMock() + return MagicMock(session=session) + + +def test_server_registers_constructor_handlers_not_decorators(): + server = mod.UnoStdioServer() + assert server.server.get_request_handler("tools/list") is not None + assert server.server.get_request_handler("tools/call") is not None + assert not hasattr(server.server, "list_tools") + assert not hasattr(server.server, "call_tool") + + +def test_list_tools_prefixes_auth_tool(monkeypatch): + gateway_tools = [ + { + "name": f"remote_{i}", + "description": "read", + "inputSchema": {"properties": {}}, + } + for i in range(7) + ] + + async def scenario(): + monkeypatch.setattr(mod.token_manager, "get_valid_token", AsyncMock(return_value=None)) + monkeypatch.setattr(mod.token_manager, "has_pending_session", lambda: False) + monkeypatch.setattr( + mod.gateway_proxy, + "list_tools", + AsyncMock(return_value={"result": {"tools": gateway_tools}}), + ) + server = mod.UnoStdioServer() + result = await server._on_list_tools(None, None) + names = [tool.name for tool in result.tools] + assert names[0] == "uno_auth" + assert names[1:] == [f"remote_{i}" for i in range(7)] + assert len(names) == 8 + assert result.tools[1].input_schema["type"] == "object" + assert "code" not in result.tools[0].input_schema["properties"] + + asyncio.run(scenario()) + + +def test_list_tools_gateway_error_returns_only_auth(monkeypatch): + async def scenario(): + monkeypatch.setattr(mod.token_manager, "get_valid_token", AsyncMock(return_value=None)) + monkeypatch.setattr( + mod.gateway_proxy, + "list_tools", + AsyncMock(side_effect=GatewayError("down")), + ) + server = mod.UnoStdioServer(link_mode=True) + result = await server._on_list_tools(None, None) + assert [tool.name for tool in result.tools] == ["uno_auth"] + assert "code" in result.tools[0].input_schema["properties"] + + asyncio.run(scenario()) + + +def test_uno_auth_status_needs_no_credentials(monkeypatch): + async def scenario(): + monkeypatch.setattr(mod.token_manager, "get_valid_token", AsyncMock(return_value=None)) + monkeypatch.setattr(mod.token_manager, "has_pending_session", lambda: False) + server = mod.UnoStdioServer() + result = await server._on_call_tool( + _ctx(), + CallToolRequestParams(name="uno_auth", arguments={"action": "status"}), + ) + body = json.loads(result.content[0].text) + assert result.is_error is False + assert body["status"] == "not_authenticated" + + asyncio.run(scenario()) + + +def test_link_mode_login_returns_auth_url_without_notifying(monkeypatch): + async def scenario(): + monkeypatch.setattr(mod.token_manager, "get_valid_token", AsyncMock(return_value=None)) + monkeypatch.setattr(mod.token_manager, "clear_pending_session", lambda: None) + monkeypatch.setattr( + mod.token_manager, + "create_link_mode_session", + AsyncMock(return_value={"auth_url": "https://example.test/oauth/start", "state": "state-id"}), + ) + server = mod.UnoStdioServer(link_mode=True) + ctx = _ctx() + result = await server._on_call_tool( + ctx, + CallToolRequestParams(name="uno_auth", arguments={"action": "login"}), + ) + body = json.loads(result.content[0].text) + assert body["status"] == "link_generated" + assert body["auth_url"] == "https://example.test/oauth/start" + ctx.session.send_tool_list_changed.assert_not_awaited() + + asyncio.run(scenario()) + + +def test_call_tool_exception_becomes_is_error(monkeypatch): + async def scenario(): + monkeypatch.setattr( + mod.token_manager, + "get_valid_token", + AsyncMock(side_effect=RuntimeError("boom")), + ) + server = mod.UnoStdioServer() + result = await server._on_call_tool( + _ctx(), + CallToolRequestParams(name="uno_auth", arguments={"action": "status"}), + ) + body = json.loads(result.content[0].text) + assert result.is_error is True + assert body["error"] == "internal_error" + + asyncio.run(scenario())