diff --git a/src/finchbot/config/env_mappings.py b/src/finchbot/config/env_mappings.py index 562f483..a224529 100644 --- a/src/finchbot/config/env_mappings.py +++ b/src/finchbot/config/env_mappings.py @@ -94,10 +94,17 @@ def get_all_mcp_env_vars() -> dict[str, Any]: servers[server_name]["args"] = [value] elif field == "url": servers[server_name]["url"] = value + elif field == "headers": + try: + headers = json.loads(value) + if isinstance(headers, dict): + servers[server_name]["headers"] = {str(k): str(v) for k, v in headers.items()} + except json.JSONDecodeError: + pass elif field == "disabled": servers[server_name]["disabled"] = value.lower() == "true" - elif field.startswith("env__"): - env_key = field[5:] + elif field == "env" and len(parts) >= 3: + env_key = parts[2] if "env" not in servers[server_name]: servers[server_name]["env"] = {} servers[server_name]["env"][env_key] = value diff --git a/src/finchbot/config/loader.py b/src/finchbot/config/loader.py index 92f1083..98d6b87 100644 --- a/src/finchbot/config/loader.py +++ b/src/finchbot/config/loader.py @@ -126,6 +126,13 @@ def _load_mcp_from_env() -> dict[str, MCPServerConfig]: servers[server_name].args = [value] elif field == "URL": servers[server_name].url = value + elif field == "HEADERS": + try: + headers = json.loads(value) + if isinstance(headers, dict): + servers[server_name].headers = {str(k): str(v) for k, v in headers.items()} + except json.JSONDecodeError: + pass elif field == "DISABLED": servers[server_name].disabled = value.lower() == "true" elif field == "ENV" and len(parts) >= 3: @@ -202,6 +209,8 @@ def load_mcp_config(workspace: Path | None = None) -> dict[str, MCPServerConfig] servers[name].args = config.args if config.url: servers[name].url = config.url + if config.headers: + servers[name].headers = config.headers else: servers[name] = config diff --git a/src/finchbot/tools/builtin/config.py b/src/finchbot/tools/builtin/config.py index a9b503e..322d8e0 100644 --- a/src/finchbot/tools/builtin/config.py +++ b/src/finchbot/tools/builtin/config.py @@ -244,7 +244,7 @@ def _add_or_update_server( command_args: list[str] | None, env: dict[str, str] | None, url: str | None, - headers: dict[str, str] | None, + headers: dict[str, str] | None = None, ) -> tuple[str, bool]: """添加或更新 MCP 服务器. @@ -558,7 +558,6 @@ async def get_mcp_tools() -> str: """ from finchbot.tools.core import ToolRegistry - workspace = _get_workspace() registry = ToolRegistry.get_instance() if not registry: @@ -573,21 +572,25 @@ async def get_mcp_tools() -> str: lines.append(f"Total: {len(mcp_tools)} tools\n") by_server: dict[str, list] = {} - for tool in mcp_tools: - server = getattr(tool, "_mcp_server_name", "unknown") + for mcp_tool in mcp_tools: + server = getattr(mcp_tool, "_mcp_server_name", "unknown") if server not in by_server: by_server[server] = [] - by_server[server].append(tool) + by_server[server].append(mcp_tool) for server_name, server_tools in sorted(by_server.items()): lines.append(f"### {server_name} ({len(server_tools)} tools)\n") - for tool in server_tools: - desc = tool.description[:150] + "..." if len(tool.description) > 150 else tool.description - lines.append(f"#### {tool.name}\n") + for server_tool in server_tools: + desc = ( + server_tool.description[:150] + "..." + if len(server_tool.description) > 150 + else server_tool.description + ) + lines.append(f"#### {server_tool.name}\n") lines.append(f"{desc}\n") - params = _get_tool_params(tool) + params = _get_tool_params(server_tool) if params: lines.append("**Parameters:**\n") for name, info in params.items(): diff --git a/src/finchbot/tools/middleware.py b/src/finchbot/tools/middleware.py index 18a5de2..edd3866 100644 --- a/src/finchbot/tools/middleware.py +++ b/src/finchbot/tools/middleware.py @@ -18,6 +18,7 @@ from pathlib import Path from typing import TYPE_CHECKING, Any +from langchain_core.tools import BaseTool from loguru import logger MIDDLEWARE_AVAILABLE = False @@ -268,12 +269,26 @@ def wrap_model_call( try: loop = asyncio.get_running_loop() - new_tools = asyncio.ensure_future(self.mcp_manager.check_and_update()) + task = loop.create_task(self.mcp_manager.check_and_update()) + + def _update_when_done(done_task: asyncio.Task) -> None: + try: + updated_tools = done_task.result() + except Exception as e: + logger.warning(f"MCP 热更新任务失败: {e}") + return + + if updated_tools is not None: + self._dynamic_tools = list(updated_tools) + logger.info(f"动态工具列表已更新: {len(updated_tools)} 个工具") + + task.add_done_callback(_update_when_done) + return handler(request) except RuntimeError: - new_tools = None + new_tools = asyncio.run(self.mcp_manager.check_and_update()) if new_tools is not None: - self._dynamic_tools = new_tools + self._dynamic_tools = list(new_tools) self._update_request_tools(request, new_tools) return handler(request) @@ -430,10 +445,22 @@ def mcp_hot_update_wrapper( import asyncio try: - asyncio.get_running_loop() - new_tools = asyncio.ensure_future(mcp_manager.check_and_update()) + loop = asyncio.get_running_loop() + + def _log_update_result(done_task: asyncio.Task) -> None: + try: + updated_tools = done_task.result() + except Exception as e: + logger.warning(f"MCP 热更新任务失败: {e}") + return + + if updated_tools is not None: + logger.info(f"MCP 热更新完成: {len(updated_tools)} 个工具") + + loop.create_task(mcp_manager.check_and_update()).add_done_callback(_log_update_result) + return handler(request) except RuntimeError: - new_tools = None + new_tools = asyncio.run(mcp_manager.check_and_update()) if new_tools is not None: existing_names = {t.name for t in request.tools} diff --git a/tests/test_config_loader.py b/tests/test_config_loader.py new file mode 100644 index 0000000..e52a3bf --- /dev/null +++ b/tests/test_config_loader.py @@ -0,0 +1,70 @@ +"""Configuration loader regression tests.""" + +from __future__ import annotations + +import json +from pathlib import Path + +from finchbot.config.env_mappings import get_all_mcp_env_vars +from finchbot.config.loader import load_mcp_config, save_mcp_config +from finchbot.config.schema import MCPServerConfig + + +def test_load_mcp_config_supports_headers_from_env(monkeypatch) -> None: + """HTTP MCP headers from env should be loaded into server config.""" + monkeypatch.setenv("FINCHBOT_MCP__REMOTE__URL", "https://example.com/mcp") + monkeypatch.setenv( + "FINCHBOT_MCP__REMOTE__HEADERS", + json.dumps({"Authorization": "Bearer env-token"}), + ) + + servers = load_mcp_config() + + assert servers["remote"].url == "https://example.com/mcp" + assert servers["remote"].headers == {"Authorization": "Bearer env-token"} + + +def test_load_mcp_config_env_headers_override_workspace_headers( + tmp_path: Path, monkeypatch +) -> None: + """Env MCP headers should override headers loaded from workspace config.""" + save_mcp_config( + { + "remote": MCPServerConfig( + url="https://example.com/mcp", + headers={"Authorization": "Bearer file-token"}, + ) + }, + tmp_path, + ) + monkeypatch.setenv( + "FINCHBOT_MCP__REMOTE__HEADERS", + json.dumps({"Authorization": "Bearer env-token"}), + ) + + servers = load_mcp_config(tmp_path) + + assert servers["remote"].headers == {"Authorization": "Bearer env-token"} + + +def test_get_all_mcp_env_vars_parses_nested_env_fields(monkeypatch) -> None: + """Nested ENV fields should be exposed by the env mapping helper.""" + monkeypatch.setenv("FINCHBOT_MCP__GITHUB__COMMAND", "npx") + monkeypatch.setenv("FINCHBOT_MCP__GITHUB__ENV__GITHUB_TOKEN", "secret") + + servers = get_all_mcp_env_vars() + + assert servers["github"]["command"] == "npx" + assert servers["github"]["env"]["GITHUB_TOKEN"] == "secret" + + +def test_get_all_mcp_env_vars_parses_headers(monkeypatch) -> None: + """HTTP MCP headers should be exposed by the env mapping helper.""" + monkeypatch.setenv( + "FINCHBOT_MCP__REMOTE__HEADERS", + json.dumps({"Authorization": "Bearer token"}), + ) + + servers = get_all_mcp_env_vars() + + assert servers["remote"]["headers"] == {"Authorization": "Bearer token"} diff --git a/tests/test_config_tools.py b/tests/test_config_tools.py index d4f296a..4b5f7ca 100644 --- a/tests/test_config_tools.py +++ b/tests/test_config_tools.py @@ -51,6 +51,28 @@ def test_add_server(self, temp_workspace: Path): assert "test-server" in data["servers"] assert data["servers"]["test-server"]["command"] == "npx" + def test_add_http_server_with_headers(self, temp_workspace: Path): + """测试添加带 headers 的 HTTP MCP 服务器.""" + result, needs_reload = config_module._add_or_update_server( + workspace=temp_workspace, + server_name="http-server", + command=None, + command_args=None, + env=None, + url="https://example.com/mcp", + headers={"Authorization": "Bearer token"}, + ) + + assert "added successfully" in result + assert needs_reload is True + + mcp_path = temp_workspace / "config" / "mcp.json" + data = json.loads(mcp_path.read_text(encoding="utf-8")) + server = data["servers"]["http-server"] + + assert server["url"] == "https://example.com/mcp" + assert server["headers"]["Authorization"] == "Bearer token" + def test_update_server(self, temp_workspace: Path): """测试更新 MCP 服务器.""" config_module._add_or_update_server( @@ -78,6 +100,37 @@ def test_update_server(self, temp_workspace: Path): data = json.loads(mcp_path.read_text(encoding="utf-8")) assert data["servers"]["test-server"]["command"] == "uvx" + def test_update_server_preserves_headers_by_default(self, temp_workspace: Path): + """测试更新 MCP 服务器时默认保留已有 headers.""" + config_module._add_or_update_server( + workspace=temp_workspace, + server_name="http-server", + command=None, + command_args=None, + env=None, + url="https://example.com/mcp", + headers={"Authorization": "Bearer token"}, + ) + + result, needs_reload = config_module._add_or_update_server( + workspace=temp_workspace, + server_name="http-server", + command=None, + command_args=None, + env={"MODE": "test"}, + url=None, + ) + + assert "updated successfully" in result + assert needs_reload is True + + mcp_path = temp_workspace / "config" / "mcp.json" + data = json.loads(mcp_path.read_text(encoding="utf-8")) + server = data["servers"]["http-server"] + + assert server["headers"]["Authorization"] == "Bearer token" + assert server["env"]["MODE"] == "test" + def test_remove_server(self, temp_workspace: Path): """测试删除 MCP 服务器.""" config_module._add_or_update_server( diff --git a/tests/test_middleware.py b/tests/test_middleware.py new file mode 100644 index 0000000..d8fc289 --- /dev/null +++ b/tests/test_middleware.py @@ -0,0 +1,86 @@ +"""Dynamic middleware regression tests.""" + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace +from typing import Any + +import pytest + +from finchbot.tools.middleware import MCPHotUpdateMiddleware + + +class DummyTool: + """Small stand-in for a LangChain tool.""" + + def __init__(self, name: str) -> None: + self.name = name + self.description = f"{name} description" + + +class DummyRegistry: + """Registry stub used by MCPHotUpdateMiddleware.""" + + def __init__(self, tools: list[DummyTool]) -> None: + self._tools = tools + + def get_tools(self) -> list[DummyTool]: + return self._tools + + +class DummyMCPManager: + """MCP manager stub returning a configured tool update.""" + + def __init__(self, tools: list[DummyTool] | None) -> None: + self.tools = tools + + async def check_and_update(self) -> list[DummyTool] | None: + await asyncio.sleep(0) + return self.tools + + +@pytest.mark.asyncio +async def test_sync_mcp_hot_update_schedules_task_without_mutating_request_to_task() -> None: + """Sync middleware should not treat an asyncio.Task as a tool list.""" + old_tool = DummyTool("old") + new_tool = DummyTool("new") + request = SimpleNamespace(tools=[old_tool]) + middleware = MCPHotUpdateMiddleware( + mcp_manager=DummyMCPManager([new_tool]), + registry=DummyRegistry([old_tool, new_tool]), + initial_tools=[old_tool], + ) + + def handler(received_request: Any) -> str: + assert received_request.tools == [old_tool] + return "ok" + + result = middleware.wrap_model_call(request, handler) + await asyncio.sleep(0.01) + + assert result == "ok" + assert middleware.tools == [new_tool] + assert all(not isinstance(tool, asyncio.Task) for tool in middleware.tools) + + +@pytest.mark.asyncio +async def test_async_mcp_hot_update_updates_current_request_tools() -> None: + """Async middleware should apply freshly loaded MCP tools immediately.""" + old_tool = DummyTool("old") + new_tool = DummyTool("new") + request = SimpleNamespace(tools=[old_tool]) + middleware = MCPHotUpdateMiddleware( + mcp_manager=DummyMCPManager([new_tool]), + registry=DummyRegistry([old_tool, new_tool]), + initial_tools=[old_tool], + ) + + async def handler(received_request: Any) -> str: + assert [tool.name for tool in received_request.tools] == ["old", "new"] + return "ok" + + result = await middleware.awrap_model_call(request, handler) + + assert result == "ok" + assert middleware.tools == [new_tool]