From b39dbedcc46307c43f8b630fa3a3baea042efb60 Mon Sep 17 00:00:00 2001 From: krypticmouse Date: Fri, 27 Mar 2026 18:09:54 +0000 Subject: [PATCH] feat: fix external MCP server integration and add streaming tool-call support MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Rebased and cleaned-up version of PR #113 by @mricharz, resolved against current main (including Codex engine, Gemini thought_signature, and agent manager fixes merged since the original PR). MCP Transport & Client: - StreamableHTTPTransport with session tracking, SSE parsing, timeouts - MCPClient.initialize() sends proper MCP handshake (protocolVersion, capabilities, clientInfo) + notifications/initialized - Fix StdioTransport constructor: command=[command] + args - MCPRequest.to_dict() with notification support (id=None) External MCP Discovery: - _discover_external_mcp supports both url (HTTP) and command (stdio) - Per-server include_tools / exclude_tools filtering - MCP clients persisted on JarvisSystem for runtime lifetime Streaming Tool-Call Support (stream_full): - StreamChunk dataclass in _stubs.py (content, tool_calls, finish_reason, usage) - Default stream_full() on InferenceEngine ABC wraps stream() for backward compat - _OpenAICompatibleEngine.stream_full() with SSE parsing - CloudEngine: _stream_full_openai (OpenAI/OpenRouter/MiniMax/Codex routing) _stream_full_anthropic (event-based → OpenAI delta format) - InstrumentedEngine, MultiEngine: stream_full delegation - GuardrailsEngine: stream_full with post-hoc security scanning (FIXED: original PR bypassed output scanning — now accumulates and scans like stream()) Other improvements: - _prepare_anthropic_messages() extracted to eliminate duplication - Default tool_choice=auto when tools are provided (OpenAI compat engines) - MCP tool injection into managed agent streaming path - Documentation: docs/user-guide/mcp-external-servers.md Tests: ~59 new tests across 8 test files, all passing. Closes PR #113 Co-Authored-By: mricharz Co-Authored-By: Claude Opus 4.6 (1M context) --- docs/user-guide/agents.md | 31 +- docs/user-guide/mcp-external-servers.md | 158 +++++ docs/user-guide/tools.md | 2 +- mkdocs.yml | 3 + src/openjarvis/engine/_openai_compat.py | 75 ++- src/openjarvis/engine/_stubs.py | 41 +- src/openjarvis/engine/cloud.py | 295 ++++++-- src/openjarvis/engine/multi.py | 32 +- src/openjarvis/mcp/__init__.py | 2 + src/openjarvis/mcp/client.py | 35 +- src/openjarvis/mcp/protocol.py | 31 +- src/openjarvis/mcp/server.py | 33 +- src/openjarvis/mcp/transport.py | 148 +++- src/openjarvis/security/guardrails.py | 77 ++- src/openjarvis/server/agent_manager_routes.py | 634 +++++++++++++----- src/openjarvis/server/stream_bridge.py | 136 ++-- src/openjarvis/system.py | 63 +- .../telemetry/instrumented_engine.py | 111 ++- tests/engine/test_cloud_stream_full.py | 466 +++++++++++++ tests/engine/test_engine_wrappers.py | 142 ++++ tests/engine/test_stream_full.py | 250 +++++++ tests/mcp/test_client_extended.py | 176 +++++ tests/mcp/test_discovery.py | 195 ++++++ tests/mcp/test_streamable_http_transport.py | 142 ++++ tests/mcp/test_transport.py | 116 +++- tests/server/test_agent_manager_routes.py | 118 ++-- tests/server/test_mcp_tools_cache.py | 154 +++++ 27 files changed, 3228 insertions(+), 438 deletions(-) create mode 100644 docs/user-guide/mcp-external-servers.md create mode 100644 tests/engine/test_cloud_stream_full.py create mode 100644 tests/engine/test_engine_wrappers.py create mode 100644 tests/engine/test_stream_full.py create mode 100644 tests/mcp/test_client_extended.py create mode 100644 tests/mcp/test_discovery.py create mode 100644 tests/mcp/test_streamable_http_transport.py create mode 100644 tests/server/test_mcp_tools_cache.py diff --git a/docs/user-guide/agents.md b/docs/user-guide/agents.md index 9f21a46f..2a988f54 100644 --- a/docs/user-guide/agents.md +++ b/docs/user-guide/agents.md @@ -626,7 +626,20 @@ These events enable the telemetry and trace systems to record detailed interacti ## Managed Agent Streaming -The Managed Agent API (`/v1/managed-agents/{id}/messages`) supports real-time SSE streaming. Send a message with `stream: true` to receive the agent's response as a Server-Sent Events stream instead of the default asynchronous queue mode. +The Managed Agent API (`/v1/managed-agents/{id}/messages`) supports **real LLM token streaming** via SSE. Send a message with `stream: true` to receive the model's response tokens as they are generated, rather than waiting for the full response. + +### How It Works + +The streaming endpoint calls `engine.stream_full()` directly, which yields `StreamChunk` objects containing content tokens, tool-call fragments, and finish reasons. This provides genuine token-by-token streaming from the LLM -- not a post-hoc word replay. + +For multi-turn tool-calling agents, the streaming loop automatically: + +1. Yields content tokens to the client as they arrive. +2. Accumulates tool-call fragments (OpenAI sends these incrementally). +3. Executes tools when `finish_reason="tool_calls"` is received. +4. Emits tool results as named SSE events (`event: tool_result`). +5. Feeds results back to the LLM for the next turn. +6. Repeats until the model produces a final text response or `max_turns` is reached. ### Streaming Messages @@ -639,19 +652,21 @@ curl -N -X POST http://localhost:8000/v1/managed-agents/{id}/messages \ The response follows the OpenAI SSE format: 1. **Content chunks** -- `data: {"choices": [{"delta": {"content": "token"}}]}` -2. **Tool results** (if the agent used tools) -- `event: tool_results\ndata: {"results": [...]}` -3. **Final chunk** -- `data: {"choices": [{"delta": {}, "finish_reason": "stop"}]}` -4. **Done sentinel** -- `data: [DONE]` +2. **Tool calls** (if the model requests tool use) -- `event: tool_calls\ndata: {"calls": [{"tool_name": "...", "arguments": "..."}]}` +3. **Tool results** -- `event: tool_result\ndata: {"tool_name": "...", "output": "..."}` +4. **Final chunk** -- `data: {"choices": [{"delta": {}, "finish_reason": "stop"}]}` +5. **Done sentinel** -- `data: [DONE]` When `stream: false` (the default), the endpoint behaves exactly as before -- the message is queued and the agent must be triggered separately via `/run`. ### Behavior Details -- The user message is always stored in the database before the agent runs. -- After streaming completes, the full agent response is persisted as an `agent_to_user` message. -- The agent is instantiated from the managed agent's stored `agent_type` and `config`. -- Conversation history from prior messages is automatically loaded as context. +- The user message is always stored in the database before streaming starts. +- After streaming completes, the full collected response is persisted as an `agent_to_user` message. +- Conversation history from prior messages is automatically loaded as LLM context. +- The engine's `stream_full()` method is used for real token streaming. Engines that do not override it fall back to the default implementation which wraps the plain `stream()` method. - If the engine is not available on the server, a `503` error is returned. +- Tool execution during streaming uses the `ToolRegistry` to find and instantiate tools. ### Python Example diff --git a/docs/user-guide/mcp-external-servers.md b/docs/user-guide/mcp-external-servers.md new file mode 100644 index 00000000..7fd316e0 --- /dev/null +++ b/docs/user-guide/mcp-external-servers.md @@ -0,0 +1,158 @@ +# External MCP Server Integration + +OpenJarvis can extend agent capabilities by connecting to external [Model Context Protocol (MCP)](https://modelcontextprotocol.io/) servers. This allows agents to use tools provided by services like Home Assistant, databases, custom APIs, or any MCP-compatible server -- without writing custom tool code. + +## How It Works + +When OpenJarvis starts, it reads the `[tools.mcp]` section in `config.toml`. For each configured server, it: + +1. Opens a connection using the appropriate transport (Streamable HTTP or stdio). +2. Performs the MCP initialize handshake (protocol version negotiation and `initialized` notification). +3. Discovers available tools via `tools/list`. +4. Wraps each discovered tool as a standard `BaseTool` so agents can call them like any built-in tool. + +If a server is unreachable or returns an error, OpenJarvis logs a warning and continues loading the remaining servers. One broken server does not prevent other tools from being available. + +## Configuration + +External MCP servers are configured in `config.toml` under `[tools.mcp]`: + +```toml +[tools.mcp] +enabled = true +servers = '[{"name": "homeassistant", "url": "http://172.16.3.1:9583/private_abc123"}]' +``` + +The `servers` value is a **JSON-encoded string** containing an array of server objects. Each object defines one external MCP server. + +!!! note + The value must be a JSON string (with single-quote TOML delimiters around it), not a native TOML array. This is because the configuration system passes it through as a single string field. + +## Server Config Schema + +Each server object supports the following fields: + +| Field | Type | Required | Description | +|------------------|----------------|----------|----------------------------------------------------------| +| `name` | string | No | Human-readable name used in log messages. Defaults to ``. | +| `url` | string | No* | URL for Streamable HTTP transport. | +| `command` | string | No* | Command to launch a stdio-based MCP server. | +| `args` | list of strings| No | Arguments passed to the stdio command. | +| `include_tools` | list of strings| No | Whitelist of tool names to import. Only these tools are loaded. | +| `exclude_tools` | list of strings| No | Blacklist of tool names to skip. All other tools are loaded. | + +*Either `url` or `command` must be provided. If neither is set, the server is skipped with a warning. + +When both `include_tools` and `exclude_tools` are specified, the whitelist is applied first, then the blacklist filters the result. + +## Examples + +### Home Assistant via Streamable HTTP + +Connect to the [ha-mcp](https://github.com/tevonsb/ha-mcp) Home Assistant add-on: + +```toml +[tools.mcp] +enabled = true +servers = '[{"name": "homeassistant", "url": "http://172.16.3.1:9583/private_abc123"}]' +``` + +This discovers all HA tools (entity control, automations, history, etc.) and makes them available to agents. + +### Stdio Server + +Launch a local MCP server as a subprocess: + +```toml +[tools.mcp] +enabled = true +servers = '[{"name": "myserver", "command": "python", "args": ["-m", "my_mcp_server"]}]' +``` + +OpenJarvis starts the process automatically, communicates via JSON-RPC over stdin/stdout, and terminates it on shutdown. + +### Multiple Servers + +```toml +[tools.mcp] +enabled = true +servers = '[{"name": "homeassistant", "url": "http://172.16.3.1:9583/private_abc123"}, {"name": "database", "command": "db-mcp-server", "args": ["--db", "postgres://localhost/mydb"]}]' +``` + +### Tool Filtering + +When a server exposes many tools but you only need a few, use `include_tools` to whitelist: + +```toml +[tools.mcp] +enabled = true +servers = '[{"name": "ha", "url": "http://172.16.3.1:9583/private_abc123", "include_tools": ["hassTurnOn", "hassTurnOff", "hassGetState"]}]' +``` + +To load everything except specific tools, use `exclude_tools`: + +```toml +[tools.mcp] +enabled = true +servers = '[{"name": "ha", "url": "http://172.16.3.1:9583/private_abc123", "exclude_tools": ["hassCreateBackup", "hassDeleteBackup"]}]' +``` + +## Transport Types + +### Streamable HTTP + +Used when the `url` field is set. The transport sends JSON-RPC requests as HTTP POST to the given URL using `httpx`. It tracks the `Mcp-Session-Id` header across requests as required by the MCP Streamable HTTP specification. + +**When to use:** Remote MCP servers, services running as HTTP endpoints (e.g., Home Assistant MCP add-on, cloud-hosted MCP servers). + +**Connection parameters:** + +- Connect timeout: 10 seconds +- Request timeout: 60 seconds + +### Stdio + +Used when the `command` field is set. OpenJarvis spawns the command as a subprocess and communicates via JSON-RPC lines on stdin/stdout. + +**When to use:** Local MCP servers distributed as CLI tools, development/testing, servers that require filesystem access on the same machine. + +!!! info "SSETransport alias" + `SSETransport` is provided as a backward-compatible alias for `StreamableHTTPTransport`. Both refer to the same implementation. + +## Error Handling + +OpenJarvis handles MCP server failures gracefully: + +- **Server unreachable:** A warning is logged and the server is skipped. All other servers and built-in tools continue to load normally. +- **Timeout:** HTTP requests time out after 60 seconds. The server is skipped with a warning. +- **Invalid config:** If the `servers` JSON is malformed or a server entry has neither `url` nor `command`, a warning is logged and that entry is skipped. +- **Tool discovery failure:** If `tools/list` fails on a server, the error is caught and the server is skipped. +- **Runtime tool call failure:** If a tool call to an external MCP server fails at runtime, it returns a `ToolResult` with `success=False` and the error message. + +No single server failure causes OpenJarvis to crash or prevents other tools from working. + +## Troubleshooting + +### Server not discovered + +1. Check that `[tools.mcp]` has `enabled = true`. +2. Verify the `servers` JSON is valid. A common mistake is using TOML arrays instead of a JSON string. +3. Check the OpenJarvis logs for warnings like `Failed to discover external MCP tools`. + +### Connection refused / timeout + +1. Verify the server is running and reachable from the OpenJarvis host: `curl -v http://host:port/`. +2. Check firewall rules between the OpenJarvis container and the MCP server. +3. For Docker deployments, ensure both containers are on the same network or use host IPs. + +### Tools not appearing + +1. Run with debug logging to see which tools were discovered. +2. Check if `include_tools` or `exclude_tools` filters are too restrictive. +3. Verify the MCP server actually exposes tools via `tools/list` (some servers only expose resources or prompts). + +### Stdio server crashes immediately + +1. Test the command manually: `python -m my_mcp_server` should start and wait for input on stdin. +2. Check stderr output in the OpenJarvis logs for error messages from the subprocess. +3. Ensure all dependencies for the MCP server are installed in the same environment. diff --git a/docs/user-guide/tools.md b/docs/user-guide/tools.md index 0e3876fd..f22bcd23 100644 --- a/docs/user-guide/tools.md +++ b/docs/user-guide/tools.md @@ -204,7 +204,7 @@ All built-in tools are registered via `@ToolRegistry.register()` and are availab | **Scheduler** | `pause_scheduled_task` | Pause an active scheduled task | | **Scheduler** | `resume_scheduled_task` | Resume a paused scheduled task | | **Scheduler** | `cancel_scheduled_task` | Cancel a scheduled task permanently | -| **Integration** | `mcp_adapter` | Bridge to external MCP tool servers | +| **Integration** | `mcp_adapter` | Bridge to external MCP tool servers (see [External MCP Servers](mcp-external-servers.md)) | --- diff --git a/mkdocs.yml b/mkdocs.yml index c8b179f9..a90a5e9c 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -164,5 +164,8 @@ nav: - Security: architecture/security.md - Design Principles: architecture/design-principles.md - API Reference: api-reference/ + - User Guide: + - Tools: user-guide/tools.md + - External MCP Servers: user-guide/mcp-external-servers.md - Leaderboard: leaderboard.md - Roadmap: development/roadmap.md diff --git a/src/openjarvis/engine/_openai_compat.py b/src/openjarvis/engine/_openai_compat.py index 1fc71e42..b9751452 100644 --- a/src/openjarvis/engine/_openai_compat.py +++ b/src/openjarvis/engine/_openai_compat.py @@ -16,6 +16,7 @@ from openjarvis.engine._base import ( estimate_prompt_tokens, messages_to_dicts, ) +from openjarvis.engine._stubs import StreamChunk logger = logging.getLogger(__name__) @@ -50,6 +51,9 @@ class _OpenAICompatibleEngine(InferenceEngine): "stream": False, **kwargs, } + # Default to tool_choice=auto when tools are provided + if "tools" in payload and "tool_choice" not in payload: + payload["tool_choice"] = "auto" try: url = f"{self._api_prefix}/chat/completions" resp = self._client.post(url, json=payload) @@ -122,6 +126,9 @@ class _OpenAICompatibleEngine(InferenceEngine): "stream": True, **kwargs, } + # Default to tool_choice=auto when tools are provided + if "tools" in payload and "tool_choice" not in payload: + payload["tool_choice"] = "auto" try: url = f"{self._api_prefix}/chat/completions" with self._client.stream("POST", url, json=payload) as resp: @@ -129,7 +136,7 @@ class _OpenAICompatibleEngine(InferenceEngine): for line in resp.iter_lines(): if not line.startswith("data:"): continue - data_str = line[len("data:"):].strip() + data_str = line[len("data:") :].strip() if data_str == "[DONE]": break try: @@ -145,16 +152,74 @@ class _OpenAICompatibleEngine(InferenceEngine): f"{self.engine_id} engine not reachable at {self._host}" ) from exc + async def stream_full( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float = 0.7, + max_tokens: int = 1024, + **kwargs: Any, + ) -> AsyncIterator["StreamChunk"]: + """Yield StreamChunks with content, tool_calls, and finish_reason.""" + msg_dicts = messages_to_dicts(messages) + payload: Dict[str, Any] = { + "model": model, + "messages": msg_dicts, + "temperature": temperature, + "max_tokens": max_tokens, + "stream": True, + **kwargs, + } + if "tools" in payload and "tool_choice" not in payload: + payload["tool_choice"] = "auto" + try: + url = f"{self._api_prefix}/chat/completions" + with self._client.stream("POST", url, json=payload) as resp: + resp.raise_for_status() + for line in resp.iter_lines(): + if not line.startswith("data:"): + continue + data_str = line[len("data:") :].strip() + if data_str == "[DONE]": + break + try: + chunk = json.loads(data_str) + except json.JSONDecodeError: + continue + choice = chunk.get("choices", [{}])[0] + delta = choice.get("delta", {}) + finish = choice.get("finish_reason") + content = delta.get("content") + tool_calls = delta.get("tool_calls") + usage = chunk.get("usage") + + if content or tool_calls or finish or usage: + yield StreamChunk( + content=content, + tool_calls=tool_calls, + finish_reason=finish, + usage=usage, + ) + except (httpx.ConnectError, httpx.TimeoutException) as exc: + raise EngineConnectionError( + f"{self.engine_id} engine not reachable at {self._host}" + ) from exc + def list_models(self) -> List[str]: try: resp = self._client.get(f"{self._api_prefix}/models") resp.raise_for_status() except ( - httpx.ConnectError, httpx.TimeoutException, httpx.HTTPStatusError, + httpx.ConnectError, + httpx.TimeoutException, + httpx.HTTPStatusError, ) as exc: logger.warning( "Failed to list models from %s at %s: %s", - self.engine_id, self._host, exc, + self.engine_id, + self._host, + exc, ) return [] data = resp.json() @@ -167,7 +232,9 @@ class _OpenAICompatibleEngine(InferenceEngine): except Exception as exc: logger.debug( "%s health check failed at %s: %s", - self.engine_id, self._host, exc, + self.engine_id, + self._host, + exc, ) return False diff --git a/src/openjarvis/engine/_stubs.py b/src/openjarvis/engine/_stubs.py index fb7feb64..04604c0e 100644 --- a/src/openjarvis/engine/_stubs.py +++ b/src/openjarvis/engine/_stubs.py @@ -14,6 +14,21 @@ from typing import Any, Dict, List, Optional, Sequence from openjarvis.core.types import Message +@dataclass(slots=True) +class StreamChunk: + """A single chunk from a streaming LLM response. + + Used by ``stream_full()`` to yield rich streaming data including + tool_calls fragments and finish_reason, unlike ``stream()`` which + only yields plain content strings. + """ + + content: Optional[str] = None + tool_calls: Optional[List[Dict[str, Any]]] = None + finish_reason: Optional[str] = None + usage: Optional[Dict[str, Any]] = None + + @dataclass(slots=True) class ResponseFormat: """Structured output configuration for inference engines. @@ -65,6 +80,30 @@ class InferenceEngine(ABC): # NOTE: must contain a yield to satisfy the type checker yield "" # pragma: no cover + async def stream_full( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float = 0.7, + max_tokens: int = 1024, + **kwargs: Any, + ) -> AsyncIterator["StreamChunk"]: + """Yield full StreamChunks including tool_calls and finish_reason. + + Default implementation wraps ``stream()`` for backward compatibility. + Engines with native tool-call streaming should override this. + """ + async for token in self.stream( + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, + ): + yield StreamChunk(content=token) + yield StreamChunk(finish_reason="stop") + @abstractmethod def list_models(self) -> List[str]: """Return identifiers of models available on this engine.""" @@ -80,4 +119,4 @@ class InferenceEngine(ABC): """Optional warm-up hook called before the first request.""" -__all__ = ["InferenceEngine", "ResponseFormat"] +__all__ = ["InferenceEngine", "ResponseFormat", "StreamChunk"] diff --git a/src/openjarvis/engine/cloud.py b/src/openjarvis/engine/cloud.py index dc9a99cd..d58b131f 100644 --- a/src/openjarvis/engine/cloud.py +++ b/src/openjarvis/engine/cloud.py @@ -6,7 +6,7 @@ import json import os import time from collections.abc import AsyncIterator, Sequence -from typing import Any, Dict, List +from typing import Any, Dict, List, Tuple import httpx @@ -17,6 +17,7 @@ from openjarvis.engine._base import ( InferenceEngine, messages_to_dicts, ) +from openjarvis.engine._stubs import StreamChunk # Pricing per million tokens (input, output) PRICING: Dict[str, tuple[float, float]] = { @@ -287,6 +288,61 @@ class CloudEngine(InferenceEngine): "url": codex_url, } + def _prepare_anthropic_messages( + self, + messages: Sequence[Message], + ) -> Tuple[str, List[Dict[str, Any]]]: + """Extract system text and convert messages to Anthropic format.""" + system_text = "" + chat_msgs: List[Dict[str, Any]] = [] + for m in messages: + if m.role.value == "system": + system_text = m.content + elif m.role.value == "tool": + tool_result_block = { + "type": "tool_result", + "tool_use_id": m.tool_call_id or "", + "content": m.content, + } + if ( + chat_msgs + and chat_msgs[-1]["role"] == "user" + and isinstance(chat_msgs[-1]["content"], list) + and chat_msgs[-1]["content"] + and chat_msgs[-1]["content"][-1].get("type") == "tool_result" + ): + chat_msgs[-1]["content"].append(tool_result_block) + else: + chat_msgs.append( + { + "role": "user", + "content": [tool_result_block], + } + ) + elif m.role.value == "assistant" and m.tool_calls: + content_blocks: List[Dict[str, Any]] = [] + if m.content: + content_blocks.append({"type": "text", "text": m.content}) + for tc in m.tool_calls: + args = tc.arguments + if isinstance(args, str): + try: + args = json.loads(args) + except (json.JSONDecodeError, TypeError): + args = {"input": args} + content_blocks.append( + { + "type": "tool_use", + "id": tc.id, + "name": tc.name, + "input": args if isinstance(args, dict) else {}, + } + ) + chat_msgs.append({"role": "assistant", "content": content_blocks}) + else: + chat_msgs.append({"role": m.role.value, "content": m.content}) + return system_text, chat_msgs + @staticmethod def _codex_build_input( messages: Sequence[Message], @@ -477,60 +533,7 @@ class CloudEngine(InferenceEngine): "ANTHROPIC_API_KEY and install " "openjarvis[inference-cloud]" ) - # Separate system message and convert to Anthropic message format - system_text = "" - chat_msgs: List[Dict[str, Any]] = [] - for m in messages: - if m.role.value == "system": - system_text = m.content - elif m.role.value == "tool": - # Anthropic expects tool results as role="user" with - # tool_result content blocks - tool_result_block = { - "type": "tool_result", - "tool_use_id": m.tool_call_id or "", - "content": m.content, - } - # Merge consecutive tool results into a single user message - if ( - chat_msgs - and chat_msgs[-1]["role"] == "user" - and isinstance(chat_msgs[-1]["content"], list) - and chat_msgs[-1]["content"] - and chat_msgs[-1]["content"][-1].get("type") == "tool_result" - ): - chat_msgs[-1]["content"].append(tool_result_block) - else: - chat_msgs.append( - { - "role": "user", - "content": [tool_result_block], - } - ) - elif m.role.value == "assistant" and m.tool_calls: - # Convert assistant messages with tool_calls to Anthropic - # content blocks (text + tool_use) - content_blocks: List[Dict[str, Any]] = [] - if m.content: - content_blocks.append({"type": "text", "text": m.content}) - for tc in m.tool_calls: - args = tc.arguments - if isinstance(args, str): - try: - args = json.loads(args) - except (json.JSONDecodeError, TypeError): - args = {"input": args} - content_blocks.append( - { - "type": "tool_use", - "id": tc.id, - "name": tc.name, - "input": args if isinstance(args, dict) else {}, - } - ) - chat_msgs.append({"role": "assistant", "content": content_blocks}) - else: - chat_msgs.append({"role": m.role.value, "content": m.content}) + system_text, chat_msgs = self._prepare_anthropic_messages(messages) create_kwargs: Dict[str, Any] = { "model": model, "messages": chat_msgs, @@ -1121,6 +1124,188 @@ class CloudEngine(InferenceEngine): if delta and delta.content: yield delta.content + # -- stream_full: rich streaming with tool_calls support ---------------- + + async def _stream_full_openai( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float, + max_tokens: int, + **kwargs: Any, + ) -> AsyncIterator[StreamChunk]: + """Yield StreamChunks from an OpenAI-compatible streaming response. + + Works for OpenAI, OpenRouter, MiniMax, and Codex. + """ + if _is_codex_model(model): + # Codex uses Responses API — fall back to base stream_full wrapper + async for chunk in super().stream_full( + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, + ): + yield chunk + return + if _is_openrouter_model(model): + client = self._openrouter_client + if client is None: + raise EngineConnectionError("OpenRouter client not available") + actual_model = model.removeprefix("openrouter/") + create_kwargs: Dict[str, Any] = { + "model": actual_model, + "messages": messages_to_dicts(messages), + "max_tokens": max_tokens, + "temperature": temperature, + "stream": True, + **kwargs, + } + elif _is_minimax_model(model): + client = self._minimax_client + if client is None: + raise EngineConnectionError("MiniMax client not available") + temperature = max(temperature, 0.01) + temperature = min(temperature, 1.0) + create_kwargs = { + "model": model, + "messages": messages_to_dicts(messages), + "max_tokens": max_tokens, + "temperature": temperature, + "stream": True, + **kwargs, + } + else: + client = self._openai_client + if client is None: + raise EngineConnectionError("OpenAI client not available") + create_kwargs = { + "model": model, + "messages": messages_to_dicts(messages), + "max_completion_tokens": max_tokens, + "stream": True, + **kwargs, + } + if not _is_openai_reasoning_model(model): + create_kwargs["temperature"] = temperature + resp = client.chat.completions.create(**create_kwargs) + for chunk in resp: + choice = chunk.choices[0] if chunk.choices else None + if not choice: + continue + delta = choice.delta + content = delta.content if delta else None + tool_calls = None + if delta and delta.tool_calls: + tool_calls = [ + { + "index": tc.index, + "id": tc.id or "", + "function": { + "name": (tc.function.name or "") if tc.function else "", + "arguments": ( + (tc.function.arguments or "") if tc.function else "" + ), + }, + } + for tc in delta.tool_calls + ] + finish = choice.finish_reason + if content or tool_calls or finish: + yield StreamChunk( + content=content, + tool_calls=tool_calls, + finish_reason=finish, + ) + + async def _stream_full_anthropic( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float, + max_tokens: int, + **kwargs: Any, + ) -> AsyncIterator[StreamChunk]: + """Yield StreamChunks from an Anthropic streaming response.""" + if self._anthropic_client is None: + raise EngineConnectionError("Anthropic client not available") + system_text, chat_msgs = self._prepare_anthropic_messages(messages) + create_kwargs: Dict[str, Any] = { + "model": model, + "messages": chat_msgs, + "temperature": temperature, + "max_tokens": max_tokens, + } + if system_text: + create_kwargs["system"] = system_text + raw_tools = kwargs.pop("tools", None) + if raw_tools: + create_kwargs["tools"] = _convert_tools_to_anthropic(raw_tools) + kwargs.pop("tool_choice", None) + + with self._anthropic_client.messages.stream(**create_kwargs) as stream: + tool_index = -1 + for event in stream: + if event.type == "content_block_start": + block = event.content_block + if block.type == "tool_use": + tool_index += 1 + yield StreamChunk( + tool_calls=[ + { + "index": tool_index, + "id": block.id, + "function": {"name": block.name, "arguments": ""}, + } + ] + ) + elif event.type == "content_block_delta": + delta = event.delta + if delta.type == "text_delta": + yield StreamChunk(content=delta.text) + elif delta.type == "input_json_delta": + yield StreamChunk( + tool_calls=[ + { + "index": tool_index, + "function": {"arguments": delta.partial_json}, + } + ] + ) + elif event.type == "message_delta": + stop_reason = event.delta.stop_reason + finish = "tool_calls" if stop_reason == "tool_use" else "stop" + yield StreamChunk(finish_reason=finish) + + async def stream_full( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float = 0.7, + max_tokens: int = 1024, + **kwargs: Any, + ) -> AsyncIterator[StreamChunk]: + """Yield StreamChunks with content, tool_calls, and finish_reason.""" + kw = dict( + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, + ) + if _is_anthropic_model(model): + async for chunk in self._stream_full_anthropic(messages, **kw): + yield chunk + elif _is_google_model(model): + async for chunk in super().stream_full(messages, **kw): + yield chunk + else: + async for chunk in self._stream_full_openai(messages, **kw): + yield chunk + def list_models(self) -> List[str]: models: List[str] = [] if self._openai_client is not None: diff --git a/src/openjarvis/engine/multi.py b/src/openjarvis/engine/multi.py index f3b9a42a..7fa48f14 100644 --- a/src/openjarvis/engine/multi.py +++ b/src/openjarvis/engine/multi.py @@ -6,7 +6,9 @@ import logging from collections.abc import AsyncIterator, Sequence from typing import Any, Dict, List +from openjarvis.core.types import Message from openjarvis.engine._base import InferenceEngine +from openjarvis.engine._stubs import StreamChunk logger = logging.getLogger(__name__) @@ -65,7 +67,7 @@ class MultiEngine(InferenceEngine): def generate( self, - messages: Sequence[Any], + messages: Sequence[Message], *, model: str, temperature: float = 0.7, @@ -73,13 +75,16 @@ class MultiEngine(InferenceEngine): **kwargs: Any, ) -> Dict[str, Any]: return self._engine_for(model).generate( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ) async def stream( self, - messages: Sequence[Any], + messages: Sequence[Message], *, model: str, temperature: float = 0.7, @@ -87,11 +92,26 @@ class MultiEngine(InferenceEngine): **kwargs: Any, ) -> AsyncIterator[str]: async for token in self._engine_for(model).stream( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ): yield token + async def stream_full( + self, + messages: Sequence[Message], + *, + model: str, + **kwargs: Any, + ) -> AsyncIterator["StreamChunk"]: + """Delegate stream_full() to the engine that owns the model.""" + engine = self._engine_for(model) + async for chunk in engine.stream_full(messages, model=model, **kwargs): + yield chunk + def list_models(self) -> List[str]: self._refresh_map() return list(self._model_map.keys()) diff --git a/src/openjarvis/mcp/__init__.py b/src/openjarvis/mcp/__init__.py index cad02d97..6fe21da9 100644 --- a/src/openjarvis/mcp/__init__.py +++ b/src/openjarvis/mcp/__init__.py @@ -8,6 +8,7 @@ from openjarvis.mcp.transport import ( MCPTransport, SSETransport, StdioTransport, + StreamableHTTPTransport, ) __all__ = [ @@ -21,4 +22,5 @@ __all__ = [ "InProcessTransport", "SSETransport", "StdioTransport", + "StreamableHTTPTransport", ] diff --git a/src/openjarvis/mcp/client.py b/src/openjarvis/mcp/client.py index 95cd7e42..c37a04a0 100644 --- a/src/openjarvis/mcp/client.py +++ b/src/openjarvis/mcp/client.py @@ -47,13 +47,36 @@ class MCPClient: def initialize(self) -> Dict[str, Any]: """Perform the MCP initialize handshake. + Sends the required client info and protocol version, then + confirms with a ``notifications/initialized`` notification + as required by the MCP specification. + Returns the server capabilities. """ - response = self._send("initialize") + params = { + "protocolVersion": "2025-03-26", + "capabilities": {}, + "clientInfo": {"name": "openjarvis", "version": "0.1.0"}, + } + response = self._send("initialize", params) self._initialized = True self._capabilities = response.result.get("capabilities", {}) + # Send the required initialized notification per MCP spec + self.notify("notifications/initialized") return response.result + def notify(self, method: str, params: Dict[str, Any] | None = None) -> None: + """Send a JSON-RPC notification (no response expected). + + Per JSON-RPC 2.0 spec, notifications omit the ``id`` field entirely. + """ + request = MCPRequest( + method=method, + params=params or {}, + id=None, # None → no id field in JSON (notification) + ) + self._transport.send_notification(request) + def list_tools(self) -> List[ToolSpec]: """Discover available tools from the server. @@ -71,7 +94,9 @@ class MCPClient: ] def call_tool( - self, name: str, arguments: Dict[str, Any] | None = None, + self, + name: str, + arguments: Dict[str, Any] | None = None, ) -> Dict[str, Any]: """Call a tool on the server. @@ -87,5 +112,11 @@ class MCPClient: """Close the transport connection.""" self._transport.close() + def __enter__(self) -> MCPClient: + return self + + def __exit__(self, *exc: Any) -> None: + self.close() + __all__ = ["MCPClient"] diff --git a/src/openjarvis/mcp/protocol.py b/src/openjarvis/mcp/protocol.py index 1b97a9ad..fe8dd99f 100644 --- a/src/openjarvis/mcp/protocol.py +++ b/src/openjarvis/mcp/protocol.py @@ -16,23 +16,34 @@ INTERNAL_ERROR = -32603 @dataclass class MCPRequest: - """JSON-RPC 2.0 request message.""" + """JSON-RPC 2.0 request message. + + Set *id* to ``None`` to create a JSON-RPC **notification** (no ``id`` + field will appear in the serialized output, and no response is expected). + """ method: str params: Dict[str, Any] = field(default_factory=dict) - id: int | str = 0 + id: Optional[int | str] = 0 jsonrpc: str = "2.0" + def to_dict(self) -> Dict[str, Any]: + """Return a dict suitable for JSON serialization. + + Omits the ``id`` key when it is ``None`` (notification). + """ + obj: Dict[str, Any] = { + "jsonrpc": self.jsonrpc, + "method": self.method, + "params": self.params, + } + if self.id is not None: + obj["id"] = self.id + return obj + def to_json(self) -> str: """Serialize to JSON string.""" - return json.dumps( - { - "jsonrpc": self.jsonrpc, - "id": self.id, - "method": self.method, - "params": self.params, - } - ) + return json.dumps(self.to_dict()) @classmethod def from_json(cls, data: str) -> MCPRequest: diff --git a/src/openjarvis/mcp/server.py b/src/openjarvis/mcp/server.py index 1988a9f4..25500372 100644 --- a/src/openjarvis/mcp/server.py +++ b/src/openjarvis/mcp/server.py @@ -70,31 +70,37 @@ class MCPServer: # Built-in API tools try: from openjarvis.tools.calculator import CalculatorTool + _tool_classes.append(CalculatorTool) except ImportError: pass try: from openjarvis.tools.think import ThinkTool + _tool_classes.append(ThinkTool) except ImportError: pass try: from openjarvis.tools.file_read import FileReadTool + _tool_classes.append(FileReadTool) except ImportError: pass try: from openjarvis.tools.web_search import WebSearchTool + _tool_classes.append(WebSearchTool) except ImportError: pass try: from openjarvis.tools.code_interpreter import CodeInterpreterTool + _tool_classes.append(CodeInterpreterTool) except ImportError: pass try: from openjarvis.tools.repl import ReplTool + _tool_classes.append(ReplTool) except ImportError: pass @@ -107,10 +113,15 @@ class MCPServer: MemorySearchTool, MemoryStoreTool, ) - _tool_classes.extend([ - MemoryStoreTool, MemoryRetrieveTool, - MemorySearchTool, MemoryIndexTool, - ]) + + _tool_classes.extend( + [ + MemoryStoreTool, + MemoryRetrieveTool, + MemorySearchTool, + MemoryIndexTool, + ] + ) except ImportError: pass @@ -121,15 +132,21 @@ class MCPServer: ChannelSendTool, ChannelStatusTool, ) - _tool_classes.extend([ - ChannelSendTool, ChannelListTool, ChannelStatusTool, - ]) + + _tool_classes.extend( + [ + ChannelSendTool, + ChannelListTool, + ChannelStatusTool, + ] + ) except ImportError: pass # LM tool (needs engine/model — instantiate with None) try: from openjarvis.tools.llm_tool import LLMTool + _tool_classes.append(LLMTool) except ImportError: pass @@ -137,6 +154,7 @@ class MCPServer: # Retrieval tool (needs backend — instantiate with None) try: from openjarvis.tools.retrieval import RetrievalTool + _tool_classes.append(RetrievalTool) except ImportError: pass @@ -150,6 +168,7 @@ class MCPServer: # Also check ToolRegistry for any user-registered tools try: from openjarvis.core.registry import ToolRegistry + known_names = {t.spec.name for t in tools} for key in ToolRegistry.keys(): if key not in known_names: diff --git a/src/openjarvis/mcp/transport.py b/src/openjarvis/mcp/transport.py index 52c7dab3..2c40bacc 100644 --- a/src/openjarvis/mcp/transport.py +++ b/src/openjarvis/mcp/transport.py @@ -2,10 +2,9 @@ from __future__ import annotations -import json import subprocess from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING, Any, List, Optional from openjarvis.mcp.protocol import MCPRequest, MCPResponse @@ -20,6 +19,15 @@ class MCPTransport(ABC): def send(self, request: MCPRequest) -> MCPResponse: """Send a request and return the response.""" + def send_notification(self, request: MCPRequest) -> None: + """Send a JSON-RPC notification (no response expected). + + The default implementation delegates to :meth:`send` and discards the + response. Transports may override this when the server returns no + body for notifications (e.g. HTTP 202 Accepted). + """ + self.send(request) + @abstractmethod def close(self) -> None: """Release transport resources.""" @@ -88,30 +96,133 @@ class StdioTransport(MCPTransport): self._process = None -class SSETransport(MCPTransport): - """JSON-RPC over HTTP with Server-Sent Events. +class StreamableHTTPTransport(MCPTransport): + """MCP Streamable HTTP transport (JSON-RPC over HTTP). - Sends requests via HTTP POST and reads SSE responses. + Uses a persistent ``httpx.Client`` session, tracks the + ``Mcp-Session-Id`` header, and sends the ``Accept`` header + required by the MCP Streamable HTTP specification. """ - def __init__(self, url: str) -> None: - self._url = url - - def send(self, request: MCPRequest) -> MCPResponse: - """Send request via HTTP POST.""" + def __init__( + self, + url: str, + *, + connect_timeout: float = 10.0, + request_timeout: float = 60.0, + ) -> None: import httpx - response = httpx.post( - self._url, - json=json.loads(request.to_json()), - headers={"Content-Type": "application/json"}, - timeout=30.0, + self._url = url + self._session_id: Optional[str] = None + self._client = httpx.Client( + timeout=httpx.Timeout( + connect=connect_timeout, + read=request_timeout, + write=request_timeout, + pool=connect_timeout, + ), ) - response.raise_for_status() - return MCPResponse.from_json(response.text) + + def _safe_url(self) -> str: + """Return scheme://host:port without path or query (avoids leaking tokens).""" + from urllib.parse import urlparse + + parsed = urlparse(self._url) + return f"{parsed.scheme}://{parsed.netloc}" + + def _build_headers(self) -> dict: + """Build common request headers.""" + headers = { + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + } + if self._session_id is not None: + headers["Mcp-Session-Id"] = self._session_id + return headers + + def _post(self, request: MCPRequest) -> Any: + """Post a request and return the raw httpx response.""" + import httpx + + headers = self._build_headers() + try: + response = self._client.post( + self._url, + json=request.to_dict(), + headers=headers, + ) + response.raise_for_status() + except httpx.ConnectError as exc: + raise RuntimeError( + f"Failed to connect to MCP server at {self._safe_url()}: {exc}" + ) from exc + except httpx.TimeoutException as exc: + raise RuntimeError( + f"Timeout communicating with MCP server at {self._safe_url()}: {exc}" + ) from exc + except httpx.HTTPStatusError as exc: + raise RuntimeError( + f"MCP server at {self._safe_url()} returned HTTP " + f"{exc.response.status_code}" + ) from exc + + # Track session id from the first response + new_session_id = response.headers.get("mcp-session-id") + if new_session_id is not None: + self._session_id = new_session_id + return response + + @staticmethod + def _extract_json_from_sse(text: str) -> str: + """Extract JSON payload from an SSE response body. + + MCP Streamable HTTP servers may respond with ``text/event-stream`` + instead of ``application/json``. In that case the body looks like:: + + event: message + data: {"jsonrpc":"2.0", ...} + + This helper finds the last ``data:`` line and returns its content, + which is the actual JSON-RPC response. + """ + last_data = "" + for line in text.splitlines(): + if line.startswith("data:"): + last_data = line[len("data:") :].strip() + if not last_data: + raise RuntimeError( + "SSE response contained no 'data:' lines" + " — cannot extract JSON-RPC payload" + ) + return last_data + + def send(self, request: MCPRequest) -> MCPResponse: + """Send request via HTTP POST following the MCP Streamable HTTP spec. + + Handles both ``application/json`` and ``text/event-stream`` responses + as allowed by the MCP Streamable HTTP specification. + """ + response = self._post(request) + content_type = response.headers.get("content-type", "") + body = response.text + if "text/event-stream" in content_type or body.lstrip().startswith("event:"): + body = self._extract_json_from_sse(body) + return MCPResponse.from_json(body) + + def send_notification(self, request: MCPRequest) -> None: + """Send a notification — accept any 2xx, don't parse the body.""" + # Track session id but don't try to parse a JSON-RPC response. + # Servers may return 202 Accepted with an empty body. + self._post(request) def close(self) -> None: - """No persistent connection to close.""" + """Close the underlying httpx client.""" + self._client.close() + + +# Backward-compatible alias +SSETransport = StreamableHTTPTransport __all__ = [ @@ -119,4 +230,5 @@ __all__ = [ "MCPTransport", "SSETransport", "StdioTransport", + "StreamableHTTPTransport", ] diff --git a/src/openjarvis/security/guardrails.py b/src/openjarvis/security/guardrails.py index aa7f4dc1..5d16883e 100644 --- a/src/openjarvis/security/guardrails.py +++ b/src/openjarvis/security/guardrails.py @@ -7,7 +7,7 @@ from typing import Any, Dict, List, Optional, Sequence from openjarvis.core.events import EventBus, EventType from openjarvis.core.types import Message -from openjarvis.engine._stubs import InferenceEngine +from openjarvis.engine._stubs import InferenceEngine, StreamChunk from openjarvis.security._stubs import BaseScanner from openjarvis.security.scanner import PIIScanner, SecretScanner from openjarvis.security.types import RedactionMode, ScanResult @@ -50,10 +50,14 @@ class GuardrailsEngine(InferenceEngine): bus: Optional[EventBus] = None, ) -> None: self._engine = engine - self._scanners: List[BaseScanner] = scanners if scanners is not None else [ - SecretScanner(), - PIIScanner(), - ] + self._scanners: List[BaseScanner] = ( + scanners + if scanners is not None + else [ + SecretScanner(), + PIIScanner(), + ] + ) self._mode = mode self._scan_input = scan_input self._scan_output = scan_output @@ -180,7 +184,9 @@ class GuardrailsEngine(InferenceEngine): processed[i] = Message( role=msg.role, content=self._handle_findings( - msg.content, result, "input", + msg.content, + result, + "input", ), name=msg.name, tool_calls=msg.tool_calls, @@ -191,8 +197,11 @@ class GuardrailsEngine(InferenceEngine): # Call wrapped engine response = self._engine.generate( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ) # Scan output @@ -219,8 +228,11 @@ class GuardrailsEngine(InferenceEngine): """Yield tokens in real-time, scan accumulated output post-hoc.""" accumulated = [] async for token in self._engine.stream( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ): accumulated.append(token) yield token @@ -248,6 +260,51 @@ class GuardrailsEngine(InferenceEngine): }, ) + async def stream_full( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float = 0.7, + max_tokens: int = 1024, + **kwargs: Any, + ) -> AsyncIterator["StreamChunk"]: + """Delegate to wrapped engine, scan accumulated output post-hoc.""" + accumulated: list[str] = [] + async for chunk in self._engine.stream_full( + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, + ): + if chunk.content: + accumulated.append(chunk.content) + yield chunk + + # Post-hoc scan of accumulated output + if self._scan_output: + full_output = "".join(accumulated) + if full_output: + result = self._scan_text(full_output) + if not result.clean and self._bus: + finding_dicts = [ + { + "pattern": f.pattern_name, + "threat": f.threat_level.value, + "description": f.description, + } + for f in result.findings + ] + self._bus.publish( + EventType.SECURITY_ALERT, + { + "direction": "output", + "findings": finding_dicts, + "mode": "stream_full_post_hoc", + }, + ) + def list_models(self) -> List[str]: """Delegate to wrapped engine.""" return self._engine.list_models() diff --git a/src/openjarvis/server/agent_manager_routes.py b/src/openjarvis/server/agent_manager_routes.py index d10da4a0..11c5404b 100644 --- a/src/openjarvis/server/agent_manager_routes.py +++ b/src/openjarvis/server/agent_manager_routes.py @@ -59,8 +59,12 @@ class FeedbackRequest(BaseModel): _BROWSER_SUB_TOOLS = { - "browser_navigate", "browser_click", "browser_type", - "browser_screenshot", "browser_extract", "browser_axtree", + "browser_navigate", + "browser_click", + "browser_type", + "browser_screenshot", + "browser_extract", + "browser_axtree", } @@ -76,7 +80,9 @@ class _LightweightSystem: def _make_lightweight_system( - engine: Any, model: str, config: Any = None, + engine: Any, + model: str, + config: Any = None, ) -> _LightweightSystem: """Build a minimal system with a plain OllamaEngine. @@ -135,9 +141,8 @@ def _ensure_registries_populated() -> None: # If registries are still empty, reload individual submodules from sys.modules if not ChannelRegistry.keys(): for mod_name in list(sys.modules): - if ( - mod_name.startswith("openjarvis.channels.") - and not mod_name.endswith("_stubs") + if mod_name.startswith("openjarvis.channels.") and not mod_name.endswith( + "_stubs" ): try: importlib.reload(sys.modules[mod_name]) @@ -192,59 +197,203 @@ def build_tools_list() -> List[Dict[str, Any]]: except Exception: spec = None cred_keys = TOOL_CREDENTIALS.get(name, []) - items.append({ - "name": name, - "description": spec.description if spec else "", - "category": spec.category if spec else "", - "source": "tool", - "requires_credentials": len(cred_keys) > 0, - "credential_keys": cred_keys, - "configured": ( - all(bool(os.environ.get(k)) for k in cred_keys) - if cred_keys else True - ), - }) + items.append( + { + "name": name, + "description": spec.description if spec else "", + "category": spec.category if spec else "", + "source": "tool", + "requires_credentials": len(cred_keys) > 0, + "credential_keys": cred_keys, + "configured": ( + all(bool(os.environ.get(k)) for k in cred_keys) + if cred_keys + else True + ), + } + ) except Exception: pass try: if any(ToolRegistry.contains(n) for n in _BROWSER_SUB_TOOLS): - items.append({ - "name": "browser", - "description": ( - "Web browser automation" - " (navigate, click, type, screenshot, extract)" - ), - "category": "browser", - "source": "tool", - "requires_credentials": False, - "credential_keys": [], - "configured": True, - }) + items.append( + { + "name": "browser", + "description": ( + "Web browser automation" + " (navigate, click, type, screenshot, extract)" + ), + "category": "browser", + "source": "tool", + "requires_credentials": False, + "credential_keys": [], + "configured": True, + } + ) except Exception: pass try: for name, _cls in ChannelRegistry.items(): cred_keys = TOOL_CREDENTIALS.get(name, []) - items.append({ - "name": name, - "description": f"{name.replace('_', ' ').title()} messaging channel", - "category": "communication", - "source": "channel", - "requires_credentials": len(cred_keys) > 0, - "credential_keys": cred_keys, - "configured": ( - all(bool(os.environ.get(k)) for k in cred_keys) - if cred_keys else True - ), - }) + items.append( + { + "name": name, + "description": f"{name.replace('_', ' ').title()} messaging channel", + "category": "communication", + "source": "channel", + "requires_credentials": len(cred_keys) > 0, + "credential_keys": cred_keys, + "configured": ( + all(bool(os.environ.get(k)) for k in cred_keys) + if cred_keys + else True + ), + } + ) except Exception: pass return items +def _merge_tool_call_fragments( + accumulated: Dict[int, Dict[str, Any]], + fragments: List[Dict[str, Any]], +) -> None: + """Merge incremental tool_call delta fragments into accumulated state. + + OpenAI-compatible APIs send tool_calls as incremental fragments keyed + by ``index``. Each fragment may contain partial ``function.name`` and/or + ``function.arguments`` strings that must be concatenated. + """ + for frag in fragments: + idx = frag.get("index", 0) + if idx not in accumulated: + accumulated[idx] = { + "id": frag.get("id", ""), + "type": "function", + "function": {"name": "", "arguments": ""}, + } + entry = accumulated[idx] + if frag.get("id"): + entry["id"] = frag["id"] + fn = frag.get("function", {}) + if fn.get("name"): + entry["function"]["name"] += fn["name"] + if fn.get("arguments"): + entry["function"]["arguments"] += fn["arguments"] + + +def _get_mcp_tools(app_state: Any) -> Tuple[List[Dict[str, Any]], Dict[str, Any]]: + """Return (openai_tools_list, mcp_adapters_by_name). + + Lazily discovers MCP tools from config and caches them on ``app_state`` + so that subsequent requests reuse the same connections. + """ + cached = getattr(app_state, "_mcp_tools_cache", None) + if cached is not None: + return cached + + import json as _json + + from openjarvis.core.config import load_config + + openai_tools: List[Dict[str, Any]] = [] + adapters_by_name: Dict[str, Any] = {} + + try: + app_config = load_config() + except Exception as exc: + logger.warning("Failed to load config for MCP discovery: %s", exc) + return openai_tools, adapters_by_name + + if not app_config.tools.mcp.enabled or not app_config.tools.mcp.servers: + return openai_tools, adapters_by_name + + from openjarvis.mcp.client import MCPClient + from openjarvis.mcp.transport import StdioTransport, StreamableHTTPTransport + from openjarvis.tools.mcp_adapter import MCPToolProvider + + # Keep clients alive so transports persist for tool calls at runtime + mcp_clients: list = getattr(app_state, "_mcp_clients", []) + + try: + server_list = _json.loads(app_config.tools.mcp.servers) + except (_json.JSONDecodeError, TypeError) as exc: + logger.warning("Failed to parse MCP server config: %s", exc) + return openai_tools, adapters_by_name + + if not isinstance(server_list, list): + return openai_tools, adapters_by_name + + for server_cfg in server_list: + cfg = _json.loads(server_cfg) if isinstance(server_cfg, str) else server_cfg + name = cfg.get("name", "") + url = cfg.get("url") + command = cfg.get("command", "") + args = cfg.get("args", []) + + try: + if url: + transport = StreamableHTTPTransport(url=url) + elif command: + transport = StdioTransport(command=[command] + args) + else: + logger.warning( + "MCP server '%s' has neither 'url' nor 'command' — skipping", + name, + ) + continue + + client = MCPClient(transport) + client.initialize() + mcp_clients.append(client) + + provider = MCPToolProvider(client) + discovered = provider.discover() + + # Per-server tool filtering + include_tools = set(cfg.get("include_tools", [])) + exclude_tools = set(cfg.get("exclude_tools", [])) + if include_tools: + discovered = [t for t in discovered if t.spec.name in include_tools] + if exclude_tools: + discovered = [t for t in discovered if t.spec.name not in exclude_tools] + + for adapter in discovered: + spec = adapter.spec + openai_tools.append( + { + "type": "function", + "function": { + "name": spec.name, + "description": spec.description, + "parameters": spec.parameters, + }, + } + ) + adapters_by_name[spec.name] = adapter + + logger.info( + "Discovered %d MCP tools from server '%s'", + len(discovered), + name, + ) + except Exception as exc: + logger.warning( + "Failed to discover MCP tools from '%s': %s", + name, + exc, + ) + + app_state._mcp_clients = mcp_clients + if openai_tools: + app_state._mcp_tools_cache = (openai_tools, adapters_by_name) + return openai_tools, adapters_by_name + + async def _stream_managed_agent( *, manager: AgentManager, @@ -253,146 +402,263 @@ async def _stream_managed_agent( message_id: str, engine: Any, bus: Any, + app_state: Any = None, ) -> StreamingResponse: - """Run a managed agent and stream the response as SSE. + """Run a managed agent with real LLM token streaming via SSE. - Instantiates the agent from its stored config, builds conversation - context from message history, executes the agent in a background - thread, and yields SSE-formatted chunks. After completion the - full response is persisted via ``manager.store_agent_response()``. + Uses ``engine.stream_full()`` to yield tokens as they arrive from the + LLM. Supports multi-turn tool-calling: when the model emits tool_calls, + they are executed and the results fed back for the next turn. """ - import asyncio import json import uuid - from openjarvis.agents._stubs import AgentContext - from openjarvis.core.registry import AgentRegistry from openjarvis.core.types import Message, Role agent_id = agent_record["id"] config = agent_record.get("config", {}) - agent_type = agent_record.get("agent_type", "orchestrator") model = config.get("model", getattr(engine, "_model", "")) + system_prompt = config.get("system_prompt") + temperature = config.get("temperature", 0.7) + max_tokens = config.get("max_tokens", 1024) + max_turns = config.get("max_turns", 10) - # Resolve the agent class from registry - agent_cls = AgentRegistry.get(agent_type) - if agent_cls is None: - # Fallback to orchestrator if the type is not registered - agent_cls = AgentRegistry.get("orchestrator") - if agent_cls is None: - raise HTTPException( - status_code=500, detail=f"Agent type '{agent_type}' not found in registry", - ) + # Build conversation messages from history + current input + llm_messages: List[Message] = [] + if system_prompt: + llm_messages.append(Message(role=Role.SYSTEM, content=system_prompt)) - # Build agent constructor kwargs from config - agent_kwargs: Dict[str, Any] = { - "engine": engine, - "model": model, - } - if bus is not None: - agent_kwargs["bus"] = bus - if config.get("system_prompt"): - agent_kwargs["system_prompt"] = config["system_prompt"] - if config.get("temperature") is not None: - agent_kwargs["temperature"] = config["temperature"] - if config.get("max_tokens") is not None: - agent_kwargs["max_tokens"] = config["max_tokens"] - if config.get("max_turns") is not None: - agent_kwargs["max_turns"] = config["max_turns"] - - try: - agent = agent_cls(**agent_kwargs) - except TypeError as exc: - logger.warning( - "Agent instantiation failed with all kwargs, retrying minimal: %s", - exc, - ) - agent = agent_cls(engine=engine, model=model) - - # Build conversation context from existing messages - ctx = AgentContext() - messages = manager.list_messages(agent_id, limit=50) - # Messages come in DESC order, reverse for chronological - for m in reversed(messages): - # Skip the message we just stored (it will be the input) + # Load prior conversation context (DESC order, reverse for chronological) + history = manager.list_messages(agent_id, limit=50) + for m in reversed(history): if m["id"] == message_id: continue if m["direction"] == "user_to_agent": - ctx.conversation.add(Message(role=Role.USER, content=m["content"])) + llm_messages.append(Message(role=Role.USER, content=m["content"])) elif m["direction"] == "agent_to_user": - ctx.conversation.add(Message(role=Role.ASSISTANT, content=m["content"])) + llm_messages.append(Message(role=Role.ASSISTANT, content=m["content"])) + + # Append the current user message + llm_messages.append(Message(role=Role.USER, content=user_content)) # Mark the user message as delivered manager.mark_message_delivered(message_id) chunk_id = f"chatcmpl-{uuid.uuid4().hex[:12]}" - async def generate(): - """Async generator yielding SSE-formatted chunks.""" - collected_content = "" + # Build extra kwargs for stream_full (e.g. tools from config) + stream_kwargs: Dict[str, Any] = {} + if config.get("tools"): + stream_kwargs["tools"] = config["tools"] - # Run agent.run() in a background thread + # Discover MCP tools and merge into stream_kwargs + mcp_adapters: Dict[str, Any] = {} + if app_state is not None: try: - result = await asyncio.to_thread(agent.run, user_content, context=ctx) + mcp_openai_tools, mcp_adapters = _get_mcp_tools(app_state) + if mcp_openai_tools: + existing_tools = stream_kwargs.get("tools", []) + stream_kwargs["tools"] = existing_tools + mcp_openai_tools + logger.info( + "Added %d MCP tools to streaming request", + len(mcp_openai_tools), + ) except Exception as exc: - logger.error("Managed agent stream error: %s", exc, exc_info=True) - error_data = { - "id": chunk_id, - "object": "chat.completion.chunk", - "model": model, - "choices": [{ - "index": 0, - "delta": {"content": f"Error: {exc}"}, - "finish_reason": "stop", - }], - } - yield f"data: {json.dumps(error_data)}\n\n" - yield "data: [DONE]\n\n" - return + logger.warning( + "Failed to get MCP tools for streaming: %s", exc, exc_info=True + ) - content = result.content or "" - collected_content = content + async def generate(): + """Async generator yielding SSE-formatted chunks with real token streaming.""" - # Emit tool results metadata if any - if result.tool_results: - tool_data = [] - for tr in result.tool_results: - tool_data.append({ - "tool_name": tr.tool_name, - "success": tr.success, - "output": tr.content, - "latency_ms": tr.latency_seconds * 1000, - }) - yield f"event: tool_results\ndata: {json.dumps({'results': tool_data})}\n\n" + collected_content = "" + messages_for_llm = list(llm_messages) + turns = 0 - # Stream content word-by-word for real-time feel - if content: - words = content.split(" ") - for i, word in enumerate(words): - token = word if i == 0 else " " + word - chunk_data = { + while turns < max_turns: + turns += 1 + turn_content = "" + tool_call_fragments: Dict[int, Dict[str, Any]] = {} + current_finish_reason = None + + try: + async for chunk in engine.stream_full( + messages_for_llm, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **stream_kwargs, + ): + # Stream content tokens immediately to the client + if chunk.content: + turn_content += chunk.content + chunk_data = { + "id": chunk_id, + "object": "chat.completion.chunk", + "model": model, + "choices": [ + { + "index": 0, + "delta": {"content": chunk.content}, + "finish_reason": None, + } + ], + } + yield f"data: {json.dumps(chunk_data)}\n\n" + + # Accumulate tool_call fragments + if chunk.tool_calls: + _merge_tool_call_fragments( + tool_call_fragments, + chunk.tool_calls, + ) + + if chunk.finish_reason: + current_finish_reason = chunk.finish_reason + + except Exception as exc: + logger.error("Managed agent stream error: %s", exc, exc_info=True) + error_data = { "id": chunk_id, "object": "chat.completion.chunk", "model": model, - "choices": [{ - "index": 0, - "delta": {"content": token}, - "finish_reason": None, - }], + "choices": [ + { + "index": 0, + "delta": {"content": f"Error: {exc}"}, + "finish_reason": "stop", + } + ], } - yield f"data: {json.dumps(chunk_data)}\n\n" - await asyncio.sleep(0.012) + yield f"data: {json.dumps(error_data)}\n\n" + yield "data: [DONE]\n\n" + return + + # Handle tool calls: execute tools and loop for next turn + if tool_call_fragments and current_finish_reason == "tool_calls": + # Build the assistant message with tool_calls + sorted_tcs = [ + tool_call_fragments[i] for i in sorted(tool_call_fragments.keys()) + ] + + # Emit tool_calls metadata as SSE event + tool_meta = [] + for tc in sorted_tcs: + tool_meta.append( + { + "tool_name": tc["function"]["name"], + "arguments": tc["function"]["arguments"], + } + ) + yield ( + f"event: tool_calls\ndata: {json.dumps({'calls': tool_meta})}\n\n" + ) + + # Add assistant message with tool_calls to conversation + from openjarvis.core.types import ToolCall as MsgToolCall + + assistant_msg = Message( + role=Role.ASSISTANT, + content=turn_content or None, + tool_calls=[ + MsgToolCall( + id=tc["id"], + name=tc["function"]["name"], + arguments=tc["function"]["arguments"], + ) + for tc in sorted_tcs + ], + ) + messages_for_llm.append(assistant_msg) + + # Execute each tool call and append results + for tc in sorted_tcs: + tool_name = tc["function"]["name"] + tool_args = tc["function"]["arguments"] + tool_result_content = f"Tool '{tool_name}' not available" + + try: + # Try MCP adapter first (external tools) + mcp_adapter = mcp_adapters.get(tool_name) + if mcp_adapter is not None: + try: + parsed_args = json.loads(tool_args) if tool_args else {} + except (json.JSONDecodeError, TypeError): + parsed_args = {} + result = mcp_adapter.execute(**parsed_args) + tool_result_content = result.content + else: + # Try to use ToolExecutor if tools are configured + from openjarvis.core.registry import ToolRegistry + from openjarvis.tools._stubs import ( + ToolCall as StubToolCall, + ) + from openjarvis.tools._stubs import ( + ToolExecutor, + ) + + tool_cls = ToolRegistry.get(tool_name) + if tool_cls is not None: + tool_instance = tool_cls() + executor = ToolExecutor(tools=[tool_instance], bus=bus) + result = executor.execute( + StubToolCall( + id=tc["id"], + name=tool_name, + arguments=tool_args, + ), + ) + tool_result_content = result.content + else: + logger.warning( + "Tool '%s' not found in registry or MCP adapters", + tool_name, + ) + except Exception as tool_exc: + logger.error( + "Tool execution error for %s: %s", + tool_name, + tool_exc, + exc_info=True, + ) + tool_result_content = f"Error executing {tool_name}: {tool_exc}" + + # Emit tool result as SSE event + tool_event_data = json.dumps( + {"tool_name": tool_name, "output": tool_result_content} + ) + yield (f"event: tool_result\ndata: {tool_event_data}\n\n") + + # Add tool result message to conversation + messages_for_llm.append( + Message( + role=Role.TOOL, + content=tool_result_content, + tool_call_id=tc["id"], + name=tool_name, + ) + ) + + # Continue to next turn (loop back to stream_full) + collected_content += turn_content + continue + + # No tool calls — this is the final response + collected_content += turn_content + break # Final chunk with finish_reason final_data = { "id": chunk_id, "object": "chat.completion.chunk", "model": model, - "choices": [{ - "index": 0, - "delta": {}, - "finish_reason": "stop", - }], + "choices": [ + { + "index": 0, + "delta": {}, + "finish_reason": "stop", + } + ], } yield f"data: {json.dumps(final_data)}\n\n" yield "data: [DONE]\n\n" @@ -403,7 +669,9 @@ async def _stream_managed_agent( manager.store_agent_response(agent_id, collected_content) except Exception as store_exc: logger.error( - "Failed to store agent response: %s", store_exc, exc_info=True, + "Failed to store agent response: %s", + store_exc, + exc_info=True, ) return StreamingResponse( @@ -507,9 +775,7 @@ def create_agent_manager_router( try: manager.start_tick(agent_id) except ValueError: - raise HTTPException( - status_code=409, detail="Agent is already running" - ) + raise HTTPException(status_code=409, detail="Agent is already running") # Re-use the server's engine + model so we don't pick a # random model from Ollama's list. @@ -523,17 +789,22 @@ def create_agent_manager_router( from openjarvis.core.events import get_event_bus executor = AgentExecutor( - manager=manager, event_bus=get_event_bus(), + manager=manager, + event_bus=get_event_bus(), ) system = _make_lightweight_system( - server_engine, server_model, server_config, + server_engine, + server_model, + server_config, ) executor.set_system(system) executor.execute_tick(agent_id) except Exception as exc: logger.error( "Run-tick failed for agent %s: %s", - agent_id, exc, exc_info=True, + agent_id, + exc, + exc_info=True, ) try: manager.end_tick(agent_id) @@ -630,7 +901,8 @@ def create_agent_manager_router( # Auto-recover error-state agents on immediate messages if req.mode == "immediate" and agent_record["status"] in ( - "error", "needs_attention", + "error", + "needs_attention", ): manager.update_agent(agent_id, status="idle") @@ -657,33 +929,39 @@ def create_agent_manager_router( def _immediate_tick(): _start = _time.time() logger.info( - "Immediate tick starting for agent %s " - "(model=%s)", - agent_id, _srv_model, + "Immediate tick starting for agent %s (model=%s)", + agent_id, + _srv_model, ) try: executor = AgentExecutor( - manager=manager, event_bus=get_event_bus(), + manager=manager, + event_bus=get_event_bus(), ) system = _make_lightweight_system( - _srv_engine, _srv_model, _srv_config, + _srv_engine, + _srv_model, + _srv_config, ) executor.set_system(system) logger.info( "Immediate tick: system ready in %.1fs, " "executing tick for agent %s", - _time.time() - _start, agent_id, + _time.time() - _start, + agent_id, ) executor.execute_tick(agent_id) logger.info( - "Immediate tick completed for agent %s " - "in %.1fs", - agent_id, _time.time() - _start, + "Immediate tick completed for agent %s in %.1fs", + agent_id, + _time.time() - _start, ) except Exception as exc: logger.error( "Immediate tick failed for agent %s: %s", - agent_id, exc, exc_info=True, + agent_id, + exc, + exc_info=True, ) try: manager.end_tick(agent_id) @@ -691,11 +969,13 @@ def create_agent_manager_router( pass manager.update_agent(agent_id, status="error") manager.update_summary_memory( - agent_id, f"ERROR: {exc}", + agent_id, + f"ERROR: {exc}", ) threading.Thread( - target=_immediate_tick, daemon=True, + target=_immediate_tick, + daemon=True, ).start() return msg @@ -715,6 +995,7 @@ def create_agent_manager_router( message_id=msg["id"], engine=engine, bus=bus, + app_state=request.app.state, ) # ── State inspection ───────────────────────────────────── @@ -820,9 +1101,7 @@ def create_agent_manager_router( @templates_router.post("/{template_id}/instantiate") async def instantiate_template(template_id: str, req: CreateAgentRequest): - return manager.create_from_template( - template_id, req.name, overrides=req.config - ) + return manager.create_from_template(template_id, req.name, overrides=req.config) # ── Global agent endpoints ─────────────────────────────── @@ -854,8 +1133,26 @@ def create_agent_manager_router( tools_router = APIRouter(prefix="/v1/tools", tags=["tools"]) @tools_router.get("") - def list_tools(): - return {"tools": build_tools_list()} + def list_tools(request: Request): + items = build_tools_list() + try: + mcp_tools, _ = _get_mcp_tools(request.app.state) + for tool in mcp_tools: + fn = tool.get("function", {}) + items.append( + { + "name": fn.get("name", ""), + "description": fn.get("description", ""), + "category": "mcp", + "source": "mcp", + "requires_credentials": False, + "credential_keys": [], + "configured": True, + } + ) + except Exception: + pass + return {"tools": items} @tools_router.post("/{tool_name}/credentials") async def save_tool_credentials(tool_name: str, request: Request): @@ -871,6 +1168,7 @@ def create_agent_manager_router( @tools_router.get("/{tool_name}/credentials/status") def credential_status(tool_name: str): from openjarvis.core.credentials import get_credential_status + return get_credential_status(tool_name) return agents_router, templates_router, global_router, tools_router diff --git a/src/openjarvis/server/stream_bridge.py b/src/openjarvis/server/stream_bridge.py index 9a96577b..386a677f 100644 --- a/src/openjarvis/server/stream_bridge.py +++ b/src/openjarvis/server/stream_bridge.py @@ -119,17 +119,15 @@ class AgentStreamBridge: from openjarvis.core.types import Message, Role for m in self._request.messages[:-1]: - role = ( - Role(m.role) - if m.role in {r.value for r in Role} - else Role.USER + role = Role(m.role) if m.role in {r.value for r in Role} else Role.USER + ctx.conversation.add( + Message( + role=role, + content=m.content or "", + name=m.name, + tool_call_id=m.tool_call_id, + ) ) - ctx.conversation.add(Message( - role=role, - content=m.content or "", - name=m.name, - tool_call_id=m.tool_call_id, - )) input_text = ( self._request.messages[-1].content if self._request.messages else "" @@ -166,9 +164,11 @@ class AgentStreamBridge: first_chunk = ChatCompletionChunk( id=self._chunk_id, model=self._model, - choices=[StreamChoice( - delta=DeltaMessage(role="assistant"), - )], + choices=[ + StreamChoice( + delta=DeltaMessage(role="assistant"), + ) + ], ) yield f"data: {first_chunk.model_dump_json()}\n\n" @@ -202,18 +202,18 @@ class AgentStreamBridge: "Please try a shorter message." ) elif "400" in error_str: - error_content = ( - f"The model returned an error: {error_str}" - ) + error_content = f"The model returned an error: {error_str}" else: error_content = f"Sorry, an error occurred: {error_str}" error_chunk = ChatCompletionChunk( id=self._chunk_id, model=self._model, - choices=[StreamChoice( - delta=DeltaMessage(content=error_content), - finish_reason="stop", - )], + choices=[ + StreamChoice( + delta=DeltaMessage(content=error_content), + finish_reason="stop", + ) + ], ) yield f"data: {error_chunk.model_dump_json()}\n\n" yield "data: [DONE]\n\n" @@ -222,31 +222,88 @@ class AgentStreamBridge: # Emit tool results metadata if any tool_results_data = [] for tr in agent_result.tool_results: - tool_results_data.append({ - "tool_name": tr.tool_name, - "success": tr.success, - "output": tr.content, - "latency_ms": tr.latency_seconds * 1000, - }) + tool_results_data.append( + { + "tool_name": tr.tool_name, + "success": tr.success, + "output": tr.content, + "latency_ms": tr.latency_seconds * 1000, + } + ) if tool_results_data: yield self._format_named_event( - "tool_results", {"results": tool_results_data}, + "tool_results", + {"results": tool_results_data}, ) - # Stream content progressively (word-by-word) for a - # real-time feel, then send a final chunk with usage. + # Stream content using real LLM token streaming via + # engine.stream_full() when the engine is available. content = agent_result.content or "" - if content: + engine = getattr(self._agent, "_engine", None) + used_real_streaming = False + + if engine is not None and hasattr(engine, "stream_full") and content: + # Re-stream using the engine for real token delivery. + # Build the same messages the agent used for its final turn. + try: + from openjarvis.core.types import Message as MsgType + from openjarvis.core.types import Role as RoleType + + replay_messages = [] + for m in self._request.messages: + role = ( + RoleType(m.role) + if m.role in {r.value for r in RoleType} + else RoleType.USER + ) + replay_messages.append( + MsgType( + role=role, + content=m.content or "", + name=m.name, + tool_call_id=m.tool_call_id, + ) + ) + + async for sc in engine.stream_full( + replay_messages, + model=self._model, + ): + if sc.content: + chunk = ChatCompletionChunk( + id=self._chunk_id, + model=self._model, + choices=[ + StreamChoice( + delta=DeltaMessage(content=sc.content), + ) + ], + ) + yield f"data: {chunk.model_dump_json()}\n\n" + used_real_streaming = True + except Exception as stream_exc: + import logging as _logging + + _logger = _logging.getLogger("openjarvis.server") + _logger.warning( + "Real streaming failed, falling back to word replay: %s", + stream_exc, + ) + + # Fallback: word-by-word replay if real streaming was not used + if not used_real_streaming and content: words = content.split(" ") for i, word in enumerate(words): token = word if i == 0 else " " + word chunk = ChatCompletionChunk( id=self._chunk_id, model=self._model, - choices=[StreamChoice( - delta=DeltaMessage(content=token), - )], + choices=[ + StreamChoice( + delta=DeltaMessage(content=token), + ) + ], ) yield f"data: {chunk.model_dump_json()}\n\n" await asyncio.sleep(0.012) @@ -254,7 +311,8 @@ class AgentStreamBridge: # Final chunk: finish_reason + usage prompt_tokens = agent_result.metadata.get("prompt_tokens", 0) completion_tokens = agent_result.metadata.get( - "completion_tokens", 0, + "completion_tokens", + 0, ) total_tokens = agent_result.metadata.get("total_tokens", 0) if total_tokens == 0: @@ -266,10 +324,12 @@ class AgentStreamBridge: final_chunk = ChatCompletionChunk( id=self._chunk_id, model=self._model, - choices=[StreamChoice( - delta=DeltaMessage(), - finish_reason="stop", - )], + choices=[ + StreamChoice( + delta=DeltaMessage(), + finish_reason="stop", + ) + ], ) final_data = json.loads(final_chunk.model_dump_json()) final_data["usage"] = UsageInfo( diff --git a/src/openjarvis/system.py b/src/openjarvis/system.py index b2dc1619..d24421f2 100644 --- a/src/openjarvis/system.py +++ b/src/openjarvis/system.py @@ -49,6 +49,7 @@ class JarvisSystem: agent_executor: Optional[Any] = None # AgentExecutor speech_backend: Optional[Any] = None # SpeechBackend _learning_orchestrator: Optional[Any] = None # LearningOrchestrator + _mcp_clients: List = field(default_factory=list) def ask( self, @@ -371,6 +372,14 @@ class JarvisSystem: channel_bridge.on_message(_on_channel_message) + def _close_mcp_clients(self) -> None: + """Close all persistent MCP client connections.""" + for client in self._mcp_clients: + try: + client.close() + except Exception: + logger.debug("Error closing MCP client", exc_info=True) + def close(self) -> None: """Release resources.""" if self.scheduler and hasattr(self.scheduler, "stop"): @@ -393,6 +402,7 @@ class JarvisSystem: self.agent_manager.close() if self.agent_scheduler is not None: self.agent_scheduler.stop() + self._close_mcp_clients() def __enter__(self) -> JarvisSystem: return self @@ -431,6 +441,7 @@ class SystemBuilder: self._workflow: Optional[bool] = None self._sessions: Optional[bool] = None self._speech: Optional[bool] = None + self._mcp_clients: List = [] def engine(self, key: str) -> SystemBuilder: self._engine_key = key @@ -671,6 +682,8 @@ class SystemBuilder: speech_backend=speech_backend, ) system._learning_orchestrator = learning_orchestrator + # Transfer MCP clients so JarvisSystem.close() can shut them down + system._mcp_clients = list(getattr(self, "_mcp_clients", [])) # Wire system reference — must happen before scheduler.start() if system.agent_executor is not None: system.agent_executor.set_system(system) @@ -1069,24 +1082,60 @@ class SystemBuilder: logger.warning("Failed to set up learning orchestrator: %s", exc) return None - @staticmethod - def _discover_external_mcp(server_cfg) -> List[BaseTool]: - """Discover tools from an external MCP server configuration.""" + def _discover_external_mcp(self, server_cfg) -> List[BaseTool]: + """Discover tools from an external MCP server configuration. + + Supports both stdio (command + args) and Streamable HTTP (url) + transports. Persists MCP clients on ``self._mcp_clients`` so + that transports stay alive for runtime tool calls. + """ import json from openjarvis.mcp.client import MCPClient - from openjarvis.mcp.transport import StdioTransport + from openjarvis.mcp.transport import StdioTransport, StreamableHTTPTransport from openjarvis.tools.mcp_adapter import MCPToolProvider cfg = json.loads(server_cfg) if isinstance(server_cfg, str) else server_cfg + name = cfg.get("name", "") + url = cfg.get("url") command = cfg.get("command", "") args = cfg.get("args", []) - if not command: + + # Build transport based on config keys + if url: + transport = StreamableHTTPTransport(url=url) + elif command: + transport = StdioTransport(command=[command] + args) + else: + logger.warning( + "MCP server '%s' has neither 'url' nor 'command' — skipping", + name, + ) return [] - transport = StdioTransport(command=command, args=args) + client = MCPClient(transport) + client.initialize() + + # Persist client so the transport stays alive for tool calls + self._mcp_clients.append(client) + provider = MCPToolProvider(client) - return provider.discover() + discovered = provider.discover() + + # Per-server tool filtering + include_tools = set(cfg.get("include_tools", [])) + exclude_tools = set(cfg.get("exclude_tools", [])) + if include_tools: + discovered = [t for t in discovered if t.spec.name in include_tools] + if exclude_tools: + discovered = [t for t in discovered if t.spec.name not in exclude_tools] + + logger.info( + "Discovered %d tools from MCP server '%s'", + len(discovered), + name, + ) + return discovered __all__ = ["JarvisSystem", "SystemBuilder"] diff --git a/src/openjarvis/telemetry/instrumented_engine.py b/src/openjarvis/telemetry/instrumented_engine.py index 1ebe5d06..29b0cb89 100644 --- a/src/openjarvis/telemetry/instrumented_engine.py +++ b/src/openjarvis/telemetry/instrumented_engine.py @@ -4,11 +4,12 @@ from __future__ import annotations import statistics import time +from collections.abc import AsyncIterator from typing import Any, Dict, List, Optional, Sequence from openjarvis.core.events import EventBus, EventType from openjarvis.core.types import Message, TelemetryRecord -from openjarvis.engine._stubs import InferenceEngine +from openjarvis.engine._stubs import InferenceEngine, StreamChunk from openjarvis.telemetry.gpu_monitor import GpuSample # --------------------------------------------------------------------------- @@ -30,8 +31,14 @@ def _percentile(data: list[float], p: float) -> float: def _compute_itl_stats(itl_values_ms: list[float]) -> dict: """Compute ITL summary statistics from a list of inter-token latencies in ms.""" if not itl_values_ms: - return {"mean": 0.0, "median": 0.0, "p90": 0.0, - "p95": 0.0, "p99": 0.0, "std": 0.0} + return { + "mean": 0.0, + "median": 0.0, + "p90": 0.0, + "p95": 0.0, + "p99": 0.0, + "std": 0.0, + } return { "mean": statistics.mean(itl_values_ms), "median": statistics.median(itl_values_ms), @@ -78,9 +85,13 @@ class InstrumentedEngine(InferenceEngine): **kwargs: Any, ) -> Dict[str, Any]: """Generate with telemetry recording.""" - self._bus.publish(EventType.INFERENCE_START, { - "model": model, "message_count": len(messages), - }) + self._bus.publish( + EventType.INFERENCE_START, + { + "model": model, + "message_count": len(messages), + }, + ) gpu_sample: Optional[GpuSample] = None energy_sample: Optional[Any] = None @@ -90,19 +101,28 @@ class InstrumentedEngine(InferenceEngine): if self._energy_monitor is not None: with self._energy_monitor.sample() as energy_sample: result = self._inner.generate( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ) elif self._gpu_monitor is not None: with self._gpu_monitor.sample() as gpu_sample: result = self._inner.generate( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ) else: result = self._inner.generate( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ) latency = time.time() - t0 @@ -155,9 +175,7 @@ class InstrumentedEngine(InferenceEngine): energy_per_output_token = ( energy_joules / completion_tokens if completion_tokens > 0 else 0.0 ) - throughput_per_watt = ( - throughput / power_watts if power_watts > 0 else 0.0 - ) + throughput_per_watt = throughput / power_watts if power_watts > 0 else 0.0 # --- Tier 2.1: Phase energy split --- decode_latency = latency - prefill_latency if prefill_latency > 0 else 0.0 @@ -171,13 +189,15 @@ class InstrumentedEngine(InferenceEngine): # --- Tier 3: Non-streaming mean ITL approximation --- mean_itl_ms = ( (decode_latency / completion_tokens) * 1000 - if completion_tokens > 0 and decode_latency > 0 else 0.0 + if completion_tokens > 0 and decode_latency > 0 + else 0.0 ) # --- Tier 4: Per-inference efficiency --- tokens_per_joule = ( completion_tokens / energy_joules - if energy_joules > 0 and completion_tokens > 0 else 0.0 + if energy_joules > 0 and completion_tokens > 0 + else 0.0 ) engine_id = getattr(self._inner, "engine_id", "unknown") @@ -275,9 +295,13 @@ class InstrumentedEngine(InferenceEngine): **kwargs: Any, ) -> Any: """Stream with per-token timing and full telemetry recording.""" - self._bus.publish(EventType.INFERENCE_START, { - "model": model, "message_count": len(messages), - }) + self._bus.publish( + EventType.INFERENCE_START, + { + "model": model, + "message_count": len(messages), + }, + ) t0 = time.time() token_timestamps: list[float] = [] @@ -289,8 +313,11 @@ class InstrumentedEngine(InferenceEngine): if self._energy_monitor is not None: with self._energy_monitor.sample() as energy_sample: async for token in self._inner.stream( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ): token_timestamps.append(time.time()) token_count += 1 @@ -298,16 +325,22 @@ class InstrumentedEngine(InferenceEngine): elif self._gpu_monitor is not None: with self._gpu_monitor.sample() as gpu_sample: async for token in self._inner.stream( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ): token_timestamps.append(time.time()) token_count += 1 yield token else: async for token in self._inner.stream( - messages, model=model, temperature=temperature, - max_tokens=max_tokens, **kwargs, + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, ): token_timestamps.append(time.time()) token_count += 1 @@ -362,9 +395,7 @@ class InstrumentedEngine(InferenceEngine): energy_per_output_token = ( energy_joules / token_count if token_count > 0 else 0.0 ) - throughput_per_watt = ( - throughput / power_watts if power_watts > 0 else 0.0 - ) + throughput_per_watt = throughput / power_watts if power_watts > 0 else 0.0 # Phase energy split decode_latency = latency - prefill_latency if prefill_latency > 0 else 0.0 @@ -378,7 +409,8 @@ class InstrumentedEngine(InferenceEngine): # Per-inference efficiency tokens_per_joule = ( token_count / energy_joules - if energy_joules > 0 and token_count > 0 else 0.0 + if energy_joules > 0 and token_count > 0 + else 0.0 ) engine_id = getattr(self._inner, "engine_id", "unknown") @@ -436,6 +468,25 @@ class InstrumentedEngine(InferenceEngine): self._bus.publish(EventType.INFERENCE_END, event_data) self._bus.publish(EventType.TELEMETRY_RECORD, {"record": record}) + async def stream_full( + self, + messages: Sequence[Message], + *, + model: str, + temperature: float = 0.7, + max_tokens: int = 1024, + **kwargs: Any, + ) -> AsyncIterator["StreamChunk"]: + """Delegate to inner engine's stream_full for tool-call support.""" + async for chunk in self._inner.stream_full( + messages, + model=model, + temperature=temperature, + max_tokens=max_tokens, + **kwargs, + ): + yield chunk + def list_models(self) -> List[str]: return self._inner.list_models() diff --git a/tests/engine/test_cloud_stream_full.py b/tests/engine/test_cloud_stream_full.py new file mode 100644 index 00000000..ce73d83a --- /dev/null +++ b/tests/engine/test_cloud_stream_full.py @@ -0,0 +1,466 @@ +"""Tests for CloudEngine.stream_full, _stream_full_openai, _stream_full_anthropic, +and _prepare_anthropic_messages.""" + +from __future__ import annotations + +from typing import Any, List +from unittest.mock import MagicMock + +import pytest + +from openjarvis.core.types import Message, Role, ToolCall +from openjarvis.engine._stubs import StreamChunk + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_cloud_engine(**overrides: Any) -> Any: + """Create a CloudEngine without calling __init__ (no env vars needed).""" + from openjarvis.engine.cloud import CloudEngine + + engine = CloudEngine.__new__(CloudEngine) + engine._openai_client = overrides.get("openai_client") + engine._anthropic_client = overrides.get("anthropic_client") + engine._google_client = overrides.get("google_client") + engine._openrouter_client = overrides.get("openrouter_client") + engine._minimax_client = overrides.get("minimax_client") + return engine + + +def _openai_chunk( + *, + content: str | None = None, + tool_calls: list | None = None, + finish_reason: str | None = None, +) -> MagicMock: + """Build a mock OpenAI streaming chunk.""" + delta = MagicMock() + delta.content = content + delta.tool_calls = tool_calls + choice = MagicMock() + choice.delta = delta + choice.finish_reason = finish_reason + chunk = MagicMock() + chunk.choices = [choice] + return chunk + + +def _openai_tool_call_delta( + *, + index: int = 0, + tc_id: str = "", + name: str = "", + arguments: str = "", +) -> MagicMock: + tc = MagicMock() + tc.index = index + tc.id = tc_id + tc.function = MagicMock() + tc.function.name = name + tc.function.arguments = arguments + return tc + + +# --------------------------------------------------------------------------- +# _stream_full_openai tests +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_stream_full_openai_content(): + """Mock OpenAI streaming response with content chunks.""" + mock_client = MagicMock() + chunks = [ + _openai_chunk(content="Hello"), + _openai_chunk(content=" world"), + _openai_chunk(finish_reason="stop"), + ] + mock_client.chat.completions.create.return_value = iter(chunks) + + engine = _make_cloud_engine(openai_client=mock_client) + msgs = [Message(role=Role.USER, content="hi")] + + result: List[StreamChunk] = [] + async for sc in engine._stream_full_openai( + msgs, + model="gpt-4o", + temperature=0.7, + max_tokens=100, + ): + result.append(sc) + + assert len(result) == 3 + assert result[0].content == "Hello" + assert result[1].content == " world" + assert result[2].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_stream_full_openai_tool_calls(): + """Mock response with tool_call deltas, verify StreamChunk.tool_calls format.""" + mock_client = MagicMock() + tc1 = _openai_tool_call_delta(index=0, tc_id="call_1", name="calc", arguments="") + tc2 = _openai_tool_call_delta(index=0, tc_id="", name="", arguments='{"x": 1}') + chunks = [ + _openai_chunk(tool_calls=[tc1]), + _openai_chunk(tool_calls=[tc2]), + _openai_chunk(finish_reason="tool_calls"), + ] + mock_client.chat.completions.create.return_value = iter(chunks) + + engine = _make_cloud_engine(openai_client=mock_client) + msgs = [Message(role=Role.USER, content="calc")] + + result: List[StreamChunk] = [] + async for sc in engine._stream_full_openai( + msgs, + model="gpt-4o", + temperature=0.7, + max_tokens=100, + ): + result.append(sc) + + assert result[0].tool_calls is not None + assert result[0].tool_calls[0]["function"]["name"] == "calc" + assert result[0].tool_calls[0]["id"] == "call_1" + assert result[1].tool_calls[0]["function"]["arguments"] == '{"x": 1}' + assert result[2].finish_reason == "tool_calls" + + +@pytest.mark.asyncio +async def test_stream_full_openai_finish_reason(): + """Verify finish_reason='tool_calls' and 'stop' propagated correctly.""" + mock_client = MagicMock() + chunks_stop = [ + _openai_chunk(content="ok"), + _openai_chunk(finish_reason="stop"), + ] + mock_client.chat.completions.create.return_value = iter(chunks_stop) + + engine = _make_cloud_engine(openai_client=mock_client) + msgs = [Message(role=Role.USER, content="hi")] + + result = [] + async for sc in engine._stream_full_openai( + msgs, + model="gpt-4o", + temperature=0.7, + max_tokens=100, + ): + result.append(sc) + + assert result[-1].finish_reason == "stop" + + # Now test tool_calls finish + tc = _openai_tool_call_delta(index=0, tc_id="c1", name="fn", arguments="{}") + chunks_tc = [ + _openai_chunk(tool_calls=[tc]), + _openai_chunk(finish_reason="tool_calls"), + ] + mock_client.chat.completions.create.return_value = iter(chunks_tc) + + result2 = [] + async for sc in engine._stream_full_openai( + msgs, + model="gpt-4o", + temperature=0.7, + max_tokens=100, + ): + result2.append(sc) + + assert result2[-1].finish_reason == "tool_calls" + + +# --------------------------------------------------------------------------- +# _stream_full_anthropic tests +# --------------------------------------------------------------------------- + + +def _anthropic_event(event_type: str, **kwargs: Any) -> MagicMock: + """Build a mock Anthropic stream event.""" + event = MagicMock() + event.type = event_type + for k, v in kwargs.items(): + setattr(event, k, v) + return event + + +@pytest.mark.asyncio +async def test_stream_full_anthropic_content(): + """Mock Anthropic stream events with text content.""" + # Build content_block_start with text type + text_block = MagicMock() + text_block.type = "text" + + # Build text delta + text_delta = MagicMock() + text_delta.type = "text_delta" + text_delta.text = "Hello world" + + # Build message_delta with stop + msg_delta = MagicMock() + msg_delta.stop_reason = "end_turn" + + events = [ + _anthropic_event("content_block_start", content_block=text_block), + _anthropic_event("content_block_delta", delta=text_delta), + _anthropic_event("message_delta", delta=msg_delta), + ] + + mock_stream = MagicMock() + mock_stream.__enter__ = MagicMock(return_value=iter(events)) + mock_stream.__exit__ = MagicMock(return_value=False) + + mock_anthropic = MagicMock() + mock_anthropic.messages.stream.return_value = mock_stream + + engine = _make_cloud_engine(anthropic_client=mock_anthropic) + msgs = [Message(role=Role.USER, content="hi")] + + result: List[StreamChunk] = [] + async for sc in engine._stream_full_anthropic( + msgs, + model="claude-sonnet-4-20250514", + temperature=0.7, + max_tokens=100, + ): + result.append(sc) + + # Should have text content and a finish reason + content_chunks = [r for r in result if r.content is not None] + assert len(content_chunks) >= 1 + assert content_chunks[0].content == "Hello world" + + finish_chunks = [r for r in result if r.finish_reason is not None] + assert len(finish_chunks) == 1 + assert finish_chunks[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_stream_full_anthropic_tool_calls(): + """Mock Anthropic tool_use events, verify OpenAI-delta-format tool_calls.""" + # content_block_start with tool_use + tool_block = MagicMock() + tool_block.type = "tool_use" + tool_block.id = "toolu_123" + tool_block.name = "get_weather" + + # input_json_delta + json_delta = MagicMock() + json_delta.type = "input_json_delta" + json_delta.partial_json = '{"city": "Berlin"}' + + # message_delta with tool_use stop + msg_delta = MagicMock() + msg_delta.stop_reason = "tool_use" + + events = [ + _anthropic_event("content_block_start", content_block=tool_block), + _anthropic_event("content_block_delta", delta=json_delta), + _anthropic_event("message_delta", delta=msg_delta), + ] + + mock_stream = MagicMock() + mock_stream.__enter__ = MagicMock(return_value=iter(events)) + mock_stream.__exit__ = MagicMock(return_value=False) + + mock_anthropic = MagicMock() + mock_anthropic.messages.stream.return_value = mock_stream + + engine = _make_cloud_engine(anthropic_client=mock_anthropic) + msgs = [Message(role=Role.USER, content="weather?")] + + result: List[StreamChunk] = [] + async for sc in engine._stream_full_anthropic( + msgs, + model="claude-sonnet-4-20250514", + temperature=0.7, + max_tokens=100, + ): + result.append(sc) + + # First chunk: tool_use start with name + assert result[0].tool_calls is not None + assert result[0].tool_calls[0]["function"]["name"] == "get_weather" + assert result[0].tool_calls[0]["id"] == "toolu_123" + + # Second chunk: arguments fragment + assert result[1].tool_calls is not None + assert result[1].tool_calls[0]["function"]["arguments"] == '{"city": "Berlin"}' + + # Third chunk: finish with tool_calls + assert result[2].finish_reason == "tool_calls" + + +@pytest.mark.asyncio +async def test_stream_full_anthropic_finish_reason(): + """message_delta with stop_reason='tool_use' maps to finish_reason='tool_calls'.""" + msg_delta_tool = MagicMock() + msg_delta_tool.stop_reason = "tool_use" + + msg_delta_stop = MagicMock() + msg_delta_stop.stop_reason = "end_turn" + + # Test tool_use -> tool_calls + events_tool = [_anthropic_event("message_delta", delta=msg_delta_tool)] + mock_stream = MagicMock() + mock_stream.__enter__ = MagicMock(return_value=iter(events_tool)) + mock_stream.__exit__ = MagicMock(return_value=False) + + mock_anthropic = MagicMock() + mock_anthropic.messages.stream.return_value = mock_stream + + engine = _make_cloud_engine(anthropic_client=mock_anthropic) + msgs = [Message(role=Role.USER, content="test")] + + result = [] + async for sc in engine._stream_full_anthropic( + msgs, + model="claude-sonnet-4-20250514", + temperature=0.7, + max_tokens=100, + ): + result.append(sc) + assert result[0].finish_reason == "tool_calls" + + # Test end_turn -> stop + events_stop = [_anthropic_event("message_delta", delta=msg_delta_stop)] + mock_stream2 = MagicMock() + mock_stream2.__enter__ = MagicMock(return_value=iter(events_stop)) + mock_stream2.__exit__ = MagicMock(return_value=False) + mock_anthropic.messages.stream.return_value = mock_stream2 + + result2 = [] + async for sc in engine._stream_full_anthropic( + msgs, + model="claude-sonnet-4-20250514", + temperature=0.7, + max_tokens=100, + ): + result2.append(sc) + assert result2[0].finish_reason == "stop" + + +# --------------------------------------------------------------------------- +# _prepare_anthropic_messages tests +# --------------------------------------------------------------------------- + + +def test_prepare_anthropic_messages_system(): + """System message extracted separately from chat messages.""" + engine = _make_cloud_engine() + msgs = [ + Message(role=Role.SYSTEM, content="You are helpful"), + Message(role=Role.USER, content="Hello"), + ] + + system_text, chat_msgs = engine._prepare_anthropic_messages(msgs) + assert system_text == "You are helpful" + assert len(chat_msgs) == 1 + assert chat_msgs[0]["role"] == "user" + assert chat_msgs[0]["content"] == "Hello" + + +def test_prepare_anthropic_messages_tool_result(): + """Tool role converted to user + tool_result content block.""" + engine = _make_cloud_engine() + msgs = [ + Message(role=Role.USER, content="What's the weather?"), + Message( + role=Role.TOOL, + content='{"temp": 20}', + tool_call_id="call_abc", + ), + ] + + system_text, chat_msgs = engine._prepare_anthropic_messages(msgs) + assert system_text == "" + assert len(chat_msgs) == 2 + # Second message is the tool result wrapped as user + tool_msg = chat_msgs[1] + assert tool_msg["role"] == "user" + assert isinstance(tool_msg["content"], list) + assert tool_msg["content"][0]["type"] == "tool_result" + assert tool_msg["content"][0]["tool_use_id"] == "call_abc" + assert tool_msg["content"][0]["content"] == '{"temp": 20}' + + +def test_prepare_anthropic_messages_tool_calls(): + """Assistant with tool_calls converted to content blocks with tool_use.""" + engine = _make_cloud_engine() + msgs = [ + Message( + role=Role.ASSISTANT, + content="Let me check.", + tool_calls=[ + ToolCall( + id="call_1", name="get_weather", arguments='{"city": "Berlin"}' + ), + ], + ), + ] + + system_text, chat_msgs = engine._prepare_anthropic_messages(msgs) + assert len(chat_msgs) == 1 + blocks = chat_msgs[0]["content"] + assert blocks[0]["type"] == "text" + assert blocks[0]["text"] == "Let me check." + assert blocks[1]["type"] == "tool_use" + assert blocks[1]["id"] == "call_1" + assert blocks[1]["name"] == "get_weather" + assert blocks[1]["input"] == {"city": "Berlin"} + + +# --------------------------------------------------------------------------- +# stream_full routing tests +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_stream_full_routes_to_anthropic(): + """model='claude-xxx' routes to _stream_full_anthropic.""" + msg_delta = MagicMock() + msg_delta.stop_reason = "end_turn" + events = [_anthropic_event("message_delta", delta=msg_delta)] + mock_stream = MagicMock() + mock_stream.__enter__ = MagicMock(return_value=iter(events)) + mock_stream.__exit__ = MagicMock(return_value=False) + + mock_anthropic = MagicMock() + mock_anthropic.messages.stream.return_value = mock_stream + + engine = _make_cloud_engine(anthropic_client=mock_anthropic) + msgs = [Message(role=Role.USER, content="test")] + + result = [] + async for sc in engine.stream_full(msgs, model="claude-sonnet-4-20250514"): + result.append(sc) + + # Verify Anthropic client was used + mock_anthropic.messages.stream.assert_called_once() + assert any(r.finish_reason is not None for r in result) + + +@pytest.mark.asyncio +async def test_stream_full_routes_to_openai(): + """model='gpt-xxx' routes to _stream_full_openai.""" + mock_client = MagicMock() + chunks = [ + _openai_chunk(content="hi"), + _openai_chunk(finish_reason="stop"), + ] + mock_client.chat.completions.create.return_value = iter(chunks) + + engine = _make_cloud_engine(openai_client=mock_client) + msgs = [Message(role=Role.USER, content="test")] + + result = [] + async for sc in engine.stream_full(msgs, model="gpt-4o"): + result.append(sc) + + # Verify OpenAI client was used + mock_client.chat.completions.create.assert_called_once() + assert result[0].content == "hi" + assert result[1].finish_reason == "stop" diff --git a/tests/engine/test_engine_wrappers.py b/tests/engine/test_engine_wrappers.py new file mode 100644 index 00000000..edfbcc4a --- /dev/null +++ b/tests/engine/test_engine_wrappers.py @@ -0,0 +1,142 @@ +"""Tests for InstrumentedEngine, GuardrailsEngine, and MultiEngine stream_full +delegation.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator +from typing import Any, Dict, List + +import pytest + +from openjarvis.core.events import EventBus +from openjarvis.core.types import Message, Role +from openjarvis.engine._stubs import InferenceEngine, StreamChunk +from openjarvis.engine.multi import MultiEngine +from openjarvis.security.guardrails import GuardrailsEngine +from openjarvis.telemetry.instrumented_engine import InstrumentedEngine + +# --------------------------------------------------------------------------- +# Fake engine that yields predetermined StreamChunks via stream_full +# --------------------------------------------------------------------------- + + +class _FakeStreamFullEngine(InferenceEngine): + engine_id = "fake-sf" + + def __init__(self, chunks: list[StreamChunk]) -> None: + self._chunks = chunks + + def generate(self, messages, *, model, **kwargs) -> Dict[str, Any]: + return {"content": "ok", "usage": {}} + + async def stream(self, messages, *, model, **kwargs) -> AsyncIterator[str]: + yield "ok" + + async def stream_full( + self, messages, *, model, **kwargs + ) -> AsyncIterator[StreamChunk]: + for c in self._chunks: + yield c + + def list_models(self) -> List[str]: + return ["fake-model"] + + def health(self) -> bool: + return True + + +# --------------------------------------------------------------------------- +# InstrumentedEngine.stream_full delegation +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_instrumented_delegates_stream_full(): + """InstrumentedEngine.stream_full delegates to inner engine.""" + expected = [ + StreamChunk(content="Hello"), + StreamChunk(content=" world"), + StreamChunk(finish_reason="stop"), + ] + inner = _FakeStreamFullEngine(expected) + bus = EventBus(record_history=True) + engine = InstrumentedEngine(inner, bus) + + result = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="fake-model", + ): + result.append(chunk) + + assert len(result) == 3 + assert result[0].content == "Hello" + assert result[1].content == " world" + assert result[2].finish_reason == "stop" + + +# --------------------------------------------------------------------------- +# GuardrailsEngine.stream_full delegation +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_guardrails_delegates_stream_full(): + """GuardrailsEngine.stream_full delegates to wrapped engine.""" + expected = [ + StreamChunk(content="safe output"), + StreamChunk(finish_reason="stop"), + ] + inner = _FakeStreamFullEngine(expected) + engine = GuardrailsEngine(inner, scanners=[]) + + result = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="fake-model", + ): + result.append(chunk) + + assert len(result) == 2 + assert result[0].content == "safe output" + assert result[1].finish_reason == "stop" + + +# --------------------------------------------------------------------------- +# MultiEngine.stream_full routing +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_multi_routes_stream_full_by_model(): + """MultiEngine routes stream_full to the correct engine by model name.""" + chunks_a = [StreamChunk(content="from A"), StreamChunk(finish_reason="stop")] + chunks_b = [StreamChunk(content="from B"), StreamChunk(finish_reason="stop")] + + engine_a = _FakeStreamFullEngine(chunks_a) + engine_a.list_models = lambda: ["model-a"] + + engine_b = _FakeStreamFullEngine(chunks_b) + engine_b.list_models = lambda: ["model-b"] + + multi = MultiEngine([("a", engine_a), ("b", engine_b)]) + + # Route to engine A + result_a = [] + async for chunk in multi.stream_full( + [Message(role=Role.USER, content="test")], + model="model-a", + ): + result_a.append(chunk) + + assert result_a[0].content == "from A" + + # Route to engine B + result_b = [] + async for chunk in multi.stream_full( + [Message(role=Role.USER, content="test")], + model="model-b", + ): + result_b.append(chunk) + + assert result_b[0].content == "from B" diff --git a/tests/engine/test_stream_full.py b/tests/engine/test_stream_full.py new file mode 100644 index 00000000..4003738c --- /dev/null +++ b/tests/engine/test_stream_full.py @@ -0,0 +1,250 @@ +"""Tests for StreamChunk dataclass and stream_full() engine method.""" + +from __future__ import annotations + +import json +from collections.abc import AsyncIterator +from typing import Any, Dict, List +from unittest.mock import MagicMock + +import pytest + +from openjarvis.core.types import Message, Role +from openjarvis.engine._stubs import InferenceEngine, StreamChunk + +# --------------------------------------------------------------------------- +# StreamChunk dataclass tests +# --------------------------------------------------------------------------- + + +class TestStreamChunk: + def test_defaults(self): + chunk = StreamChunk() + assert chunk.content is None + assert chunk.tool_calls is None + assert chunk.finish_reason is None + assert chunk.usage is None + + def test_content_only(self): + chunk = StreamChunk(content="hello") + assert chunk.content == "hello" + assert chunk.finish_reason is None + + def test_finish_reason(self): + chunk = StreamChunk(finish_reason="stop") + assert chunk.content is None + assert chunk.finish_reason == "stop" + + def test_tool_calls(self): + tc = [{"index": 0, "function": {"name": "calc", "arguments": "{}"}}] + chunk = StreamChunk(tool_calls=tc) + assert chunk.tool_calls == tc + + def test_usage(self): + usage = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} + chunk = StreamChunk(usage=usage) + assert chunk.usage == usage + + def test_all_fields(self): + chunk = StreamChunk( + content="hi", + tool_calls=[{"index": 0}], + finish_reason="tool_calls", + usage={"total_tokens": 1}, + ) + assert chunk.content == "hi" + assert chunk.tool_calls is not None + assert chunk.finish_reason == "tool_calls" + assert chunk.usage is not None + + +# --------------------------------------------------------------------------- +# Concrete engine stub for testing default stream_full() +# --------------------------------------------------------------------------- + + +class _FakeEngine(InferenceEngine): + """Minimal engine that yields predefined tokens via stream().""" + + engine_id = "fake" + + def __init__(self, tokens: list[str]) -> None: + self._tokens = tokens + + def generate(self, messages, *, model, **kwargs) -> Dict[str, Any]: + return {"content": "".join(self._tokens), "usage": {}} + + async def stream(self, messages, *, model, **kwargs) -> AsyncIterator[str]: + for t in self._tokens: + yield t + + def list_models(self) -> List[str]: + return ["fake-model"] + + def health(self) -> bool: + return True + + +class TestDefaultStreamFull: + """Test the default stream_full() implementation that wraps stream().""" + + @pytest.mark.asyncio + async def test_wraps_stream_tokens(self): + engine = _FakeEngine(["Hello", " world", "!"]) + chunks = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="fake-model", + ): + chunks.append(chunk) + + # Should have 3 content chunks + 1 finish chunk + assert len(chunks) == 4 + assert chunks[0].content == "Hello" + assert chunks[1].content == " world" + assert chunks[2].content == "!" + assert chunks[3].finish_reason == "stop" + assert chunks[3].content is None + + @pytest.mark.asyncio + async def test_empty_stream(self): + engine = _FakeEngine([]) + chunks = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="fake-model", + ): + chunks.append(chunk) + + # Should have just the finish chunk + assert len(chunks) == 1 + assert chunks[0].finish_reason == "stop" + + @pytest.mark.asyncio + async def test_kwargs_passed_through(self): + """Verify that temperature/max_tokens reach stream().""" + engine = _FakeEngine(["ok"]) + chunks = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="fake-model", + temperature=0.1, + max_tokens=50, + ): + chunks.append(chunk) + + assert len(chunks) == 2 + assert chunks[0].content == "ok" + + +# --------------------------------------------------------------------------- +# OpenAI-compatible stream_full() with mock HTTP response +# --------------------------------------------------------------------------- + + +class TestOpenAICompatStreamFull: + """Test _OpenAICompatibleEngine.stream_full() with mocked HTTP.""" + + @pytest.mark.asyncio + async def test_parses_sse_with_content_and_finish(self): + from openjarvis.engine._openai_compat import _OpenAICompatibleEngine + + # Build mock SSE lines + sse_lines = [] + for token in ["Hello", " world"]: + chunk = { + "choices": [{"delta": {"content": token}, "finish_reason": None}], + } + sse_lines.append(f"data: {json.dumps(chunk)}") + # Final chunk with finish_reason + final = { + "choices": [{"delta": {}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, + } + sse_lines.append(f"data: {json.dumps(final)}") + sse_lines.append("data: [DONE]") + + # Mock the httpx client stream context manager + mock_resp = MagicMock() + mock_resp.raise_for_status = MagicMock() + mock_resp.iter_lines.return_value = iter(sse_lines) + + engine = _OpenAICompatibleEngine.__new__(_OpenAICompatibleEngine) + engine.engine_id = "test" + engine._host = "http://localhost:8000" + engine._api_prefix = "/v1" + + mock_client = MagicMock() + mock_stream_ctx = MagicMock() + mock_stream_ctx.__enter__ = MagicMock(return_value=mock_resp) + mock_stream_ctx.__exit__ = MagicMock(return_value=False) + mock_client.stream.return_value = mock_stream_ctx + engine._client = mock_client + + chunks = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="test-model", + ): + chunks.append(chunk) + + # Should have: "Hello", " world", finish+usage + assert len(chunks) == 3 + assert chunks[0].content == "Hello" + assert chunks[1].content == " world" + assert chunks[2].finish_reason == "stop" + assert chunks[2].usage is not None + assert chunks[2].usage["total_tokens"] == 7 + + @pytest.mark.asyncio + async def test_parses_tool_call_fragments(self): + from openjarvis.engine._openai_compat import _OpenAICompatibleEngine + + # Simulate streamed tool_call fragments + _tc1 = ( + '{"choices": [{"delta": {"tool_calls": [{"index": 0, "id": "call_1",' + ' "function": {"name": "calc", "arguments": ""}}]},' + ' "finish_reason": null}]}' + ) + _tc2 = ( + '{"choices": [{"delta": {"tool_calls": [{"index": 0,' + ' "function": {"name": "", "arguments": "{\\"x\\": 1}"}}]},' + ' "finish_reason": null}]}' + ) + sse_lines = [ + f"data: {_tc1}", + f"data: {_tc2}", + 'data: {"choices": [{"delta": {}, "finish_reason": "tool_calls"}]}', + "data: [DONE]", + ] + + mock_resp = MagicMock() + mock_resp.raise_for_status = MagicMock() + mock_resp.iter_lines.return_value = iter(sse_lines) + + engine = _OpenAICompatibleEngine.__new__(_OpenAICompatibleEngine) + engine.engine_id = "test" + engine._host = "http://localhost:8000" + engine._api_prefix = "/v1" + + mock_client = MagicMock() + mock_stream_ctx = MagicMock() + mock_stream_ctx.__enter__ = MagicMock(return_value=mock_resp) + mock_stream_ctx.__exit__ = MagicMock(return_value=False) + mock_client.stream.return_value = mock_stream_ctx + engine._client = mock_client + + chunks = [] + async for chunk in engine.stream_full( + [Message(role=Role.USER, content="test")], + model="test-model", + ): + chunks.append(chunk) + + # First chunk has tool_calls with name + assert chunks[0].tool_calls is not None + assert chunks[0].tool_calls[0]["function"]["name"] == "calc" + # Second chunk has arguments fragment + assert chunks[1].tool_calls is not None + # Third chunk has finish_reason="tool_calls" + assert chunks[2].finish_reason == "tool_calls" diff --git a/tests/mcp/test_client_extended.py b/tests/mcp/test_client_extended.py new file mode 100644 index 00000000..5857a5e1 --- /dev/null +++ b/tests/mcp/test_client_extended.py @@ -0,0 +1,176 @@ +"""Extended tests for MCPClient — initialize params, notify, context manager.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from openjarvis.mcp.client import MCPClient +from openjarvis.mcp.protocol import MCPRequest, MCPResponse +from openjarvis.tools._stubs import ToolSpec + + +@pytest.fixture +def mock_transport(): + """A mock transport that returns configurable responses.""" + transport = MagicMock() + transport.send.return_value = MCPResponse(result={}) + return transport + + +class TestInitialize: + def test_sends_correct_params(self, mock_transport): + """initialize() must send protocolVersion, capabilities, clientInfo.""" + mock_transport.send.return_value = MCPResponse( + result={"protocolVersion": "2025-03-26", "capabilities": {"tools": {}}} + ) + client = MCPClient(mock_transport) + client.initialize() + + # First call is the initialize request + init_call = mock_transport.send.call_args_list[0] + req = init_call[0][0] + assert isinstance(req, MCPRequest) + assert req.method == "initialize" + assert req.params["protocolVersion"] == "2025-03-26" + assert req.params["capabilities"] == {} + assert req.params["clientInfo"]["name"] == "openjarvis" + assert req.params["clientInfo"]["version"] == "0.1.0" + + def test_sends_initialized_notification(self, mock_transport): + """After initialize, notifications/initialized must be sent via + send_notification.""" + mock_transport.send.return_value = MCPResponse( + result={"protocolVersion": "2025-03-26", "capabilities": {}} + ) + client = MCPClient(mock_transport) + client.initialize() + + # initialize request goes via send(), notification via send_notification() + assert mock_transport.send.call_count == 1 + mock_transport.send_notification.assert_called_once() + req = mock_transport.send_notification.call_args[0][0] + assert req.method == "notifications/initialized" + assert req.id is None # notifications must not have an id + + def test_stores_capabilities(self, mock_transport): + """initialize() should store server capabilities.""" + mock_transport.send.return_value = MCPResponse( + result={"capabilities": {"tools": {"listChanged": True}}} + ) + client = MCPClient(mock_transport) + client.initialize() + assert client._capabilities == {"tools": {"listChanged": True}} + + +class TestNotify: + def test_notify_sends_request(self, mock_transport): + """notify() should send a notification with the given method and params.""" + client = MCPClient(mock_transport) + client.notify("notifications/cancelled", {"requestId": 42}) + + mock_transport.send_notification.assert_called_once() + req = mock_transport.send_notification.call_args[0][0] + assert req.method == "notifications/cancelled" + assert req.params == {"requestId": 42} + assert req.id is None # notifications must omit id + + def test_notify_defaults_empty_params(self, mock_transport): + """notify() with no params should send empty dict.""" + client = MCPClient(mock_transport) + client.notify("notifications/initialized") + + req = mock_transport.send_notification.call_args[0][0] + assert req.params == {} + + def test_notify_json_has_no_id(self, mock_transport): + """The serialized notification JSON must not contain an 'id' field.""" + import json + + client = MCPClient(mock_transport) + client.notify("notifications/initialized") + + req = mock_transport.send_notification.call_args[0][0] + payload = json.loads(req.to_json()) + assert "id" not in payload + + +class TestContextManager: + def test_context_manager_calls_close(self, mock_transport): + """Using MCPClient as context manager should call close on exit.""" + with MCPClient(mock_transport) as client: + assert client is not None + mock_transport.close.assert_called_once() + + def test_context_manager_closes_on_exception(self, mock_transport): + """close() should be called even if an exception occurs.""" + with pytest.raises(ValueError): + with MCPClient(mock_transport): + raise ValueError("test error") + mock_transport.close.assert_called_once() + + +class TestListTools: + def test_parses_tool_specs(self, mock_transport): + """list_tools() should return ToolSpec objects from server response.""" + mock_transport.send.return_value = MCPResponse( + result={ + "tools": [ + { + "name": "get_entities", + "description": "Get HA entities", + "inputSchema": { + "type": "object", + "properties": {"domain": {"type": "string"}}, + }, + }, + { + "name": "call_service", + "description": "Call HA service", + "inputSchema": {"type": "object", "properties": {}}, + }, + ] + } + ) + client = MCPClient(mock_transport) + tools = client.list_tools() + + assert len(tools) == 2 + assert all(isinstance(t, ToolSpec) for t in tools) + assert tools[0].name == "get_entities" + assert tools[0].description == "Get HA entities" + assert "properties" in tools[0].parameters + assert tools[1].name == "call_service" + + def test_empty_tools_list(self, mock_transport): + """list_tools() with no tools should return empty list.""" + mock_transport.send.return_value = MCPResponse(result={"tools": []}) + client = MCPClient(mock_transport) + assert client.list_tools() == [] + + +class TestCallTool: + def test_call_tool_sends_correct_params(self, mock_transport): + """call_tool() should send method=tools/call with name and arguments.""" + mock_transport.send.return_value = MCPResponse( + result={"content": [{"type": "text", "text": "ok"}], "isError": False} + ) + client = MCPClient(mock_transport) + result = client.call_tool("get_entities", {"domain": "light"}) + + req = mock_transport.send.call_args[0][0] + assert req.method == "tools/call" + assert req.params == {"name": "get_entities", "arguments": {"domain": "light"}} + assert result["isError"] is False + + def test_call_tool_no_arguments(self, mock_transport): + """call_tool() with no arguments passes empty dict.""" + mock_transport.send.return_value = MCPResponse( + result={"content": [{"type": "text", "text": "done"}], "isError": False} + ) + client = MCPClient(mock_transport) + client.call_tool("ping") + + req = mock_transport.send.call_args[0][0] + assert req.params == {"name": "ping", "arguments": {}} diff --git a/tests/mcp/test_discovery.py b/tests/mcp/test_discovery.py new file mode 100644 index 00000000..ac9c061e --- /dev/null +++ b/tests/mcp/test_discovery.py @@ -0,0 +1,195 @@ +"""Tests for _discover_external_mcp in SystemBuilder.""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from openjarvis.tools._stubs import ToolSpec + + +def _make_mock_tool(name: str) -> MagicMock: + """Create a mock BaseTool with the given name.""" + tool = MagicMock() + tool.spec = ToolSpec(name=name, description=f"Mock {name}") + return tool + + +@pytest.fixture +def builder(): + """Create a minimal SystemBuilder instance for testing _discover_external_mcp.""" + from openjarvis.system import SystemBuilder + + def _minimal_init(self): + self._mcp_clients = [] + + with patch.object(SystemBuilder, "__init__", _minimal_init): + instance = SystemBuilder.__new__(SystemBuilder) + instance.__init__() + return instance + + +# Patch targets: the method uses local imports, so we patch the actual classes +# in their source modules (which is where the local from-imports resolve to). +_PATCH_HTTP = "openjarvis.mcp.transport.StreamableHTTPTransport" +_PATCH_STDIO = "openjarvis.mcp.transport.StdioTransport" +_PATCH_CLIENT = "openjarvis.mcp.client.MCPClient" +_PATCH_PROVIDER = "openjarvis.tools.mcp_adapter.MCPToolProvider" +_PATCH_LOGGER = "openjarvis.system.logger" + + +class TestDiscoverHTTPServer: + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_HTTP) + def test_url_config_uses_http_transport( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """Config with 'url' should create StreamableHTTPTransport.""" + mock_tools = [_make_mock_tool("get_entities"), _make_mock_tool("call_service")] + mock_provider_cls.return_value.discover.return_value = mock_tools + + cfg = {"name": "ha-mcp", "url": "http://172.16.3.1:9583/mcp"} + result = builder._discover_external_mcp(cfg) + + mock_transport_cls.assert_called_once_with(url="http://172.16.3.1:9583/mcp") + mock_client_cls.return_value.initialize.assert_called_once() + assert len(result) == 2 + assert result[0].spec.name == "get_entities" + + +class TestDiscoverStdioServer: + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_STDIO) + def test_command_config_uses_stdio_transport( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """Config with 'command' + 'args' should create StdioTransport.""" + mock_tools = [_make_mock_tool("read_file")] + mock_provider_cls.return_value.discover.return_value = mock_tools + + cfg = {"name": "fs-server", "command": "node", "args": ["server.js", "--stdio"]} + result = builder._discover_external_mcp(cfg) + + mock_transport_cls.assert_called_once_with( + command=["node", "server.js", "--stdio"] + ) + assert len(result) == 1 + assert result[0].spec.name == "read_file" + + +class TestDiscoverInvalidConfig: + @patch(_PATCH_LOGGER) + def test_no_url_no_command_returns_empty(self, mock_logger, builder): + """Config with neither 'url' nor 'command' should return [] and log warning.""" + cfg = {"name": "broken-server"} + result = builder._discover_external_mcp(cfg) + + assert result == [] + mock_logger.warning.assert_called_once() + assert "neither" in mock_logger.warning.call_args[0][0].lower() + + +class TestToolFiltering: + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_HTTP) + def test_include_tools_filter( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """include_tools should keep only the listed tools.""" + mock_tools = [ + _make_mock_tool("tool1"), + _make_mock_tool("tool2"), + _make_mock_tool("tool3"), + ] + mock_provider_cls.return_value.discover.return_value = mock_tools + + cfg = { + "name": "filtered", + "url": "http://localhost:8080/mcp", + "include_tools": ["tool1"], + } + result = builder._discover_external_mcp(cfg) + + assert len(result) == 1 + assert result[0].spec.name == "tool1" + + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_HTTP) + def test_exclude_tools_filter( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """exclude_tools should remove the listed tools.""" + mock_tools = [ + _make_mock_tool("tool1"), + _make_mock_tool("tool2"), + _make_mock_tool("tool3"), + ] + mock_provider_cls.return_value.discover.return_value = mock_tools + + cfg = { + "name": "filtered", + "url": "http://localhost:8080/mcp", + "exclude_tools": ["tool2"], + } + result = builder._discover_external_mcp(cfg) + + names = [t.spec.name for t in result] + assert "tool2" not in names + assert "tool1" in names + assert "tool3" in names + + +class TestClientPersistence: + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_HTTP) + def test_client_stored_in_mcp_clients( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """After discovery, the MCPClient should be persisted on _mcp_clients.""" + mock_provider_cls.return_value.discover.return_value = [] + + cfg = {"name": "test", "url": "http://localhost:8080/mcp"} + builder._discover_external_mcp(cfg) + + assert hasattr(builder, "_mcp_clients") + assert len(builder._mcp_clients) == 1 + assert builder._mcp_clients[0] is mock_client_cls.return_value + + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_HTTP) + def test_multiple_servers_accumulate_clients( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """Multiple discover calls should accumulate clients.""" + mock_provider_cls.return_value.discover.return_value = [] + + for i in range(3): + cfg = {"name": f"server-{i}", "url": f"http://localhost:{8080 + i}/mcp"} + builder._discover_external_mcp(cfg) + + assert len(builder._mcp_clients) == 3 + + +class TestStringConfig: + @patch(_PATCH_PROVIDER) + @patch(_PATCH_CLIENT) + @patch(_PATCH_HTTP) + def test_json_string_config_parsed( + self, mock_transport_cls, mock_client_cls, mock_provider_cls, builder + ): + """Config passed as JSON string should be parsed correctly.""" + import json + + mock_provider_cls.return_value.discover.return_value = [] + + cfg_str = json.dumps({"name": "test", "url": "http://localhost:8080/mcp"}) + builder._discover_external_mcp(cfg_str) + + mock_transport_cls.assert_called_once_with(url="http://localhost:8080/mcp") diff --git a/tests/mcp/test_streamable_http_transport.py b/tests/mcp/test_streamable_http_transport.py new file mode 100644 index 00000000..0bb9b979 --- /dev/null +++ b/tests/mcp/test_streamable_http_transport.py @@ -0,0 +1,142 @@ +"""Tests for the StreamableHTTPTransport class.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import pytest + +from openjarvis.mcp.protocol import MCPRequest + + +@pytest.fixture +def _mock_httpx_client(): + """Patch httpx.Client so no real HTTP connections are made.""" + with patch("httpx.Client") as mock_cls, patch("httpx.Timeout"): + mock_instance = MagicMock() + mock_cls.return_value = mock_instance + yield mock_instance + + +def _make_http_response(result, *, session_id=None): + """Build a mock httpx.Response with the given JSON-RPC result.""" + resp = MagicMock() + resp.text = json.dumps({"jsonrpc": "2.0", "id": 1, "result": result}) + resp.raise_for_status = MagicMock() + headers = {} + if session_id is not None: + headers["mcp-session-id"] = session_id + resp.headers = headers + return resp + + +class TestStreamableHTTPTransport: + def test_send_request(self, _mock_httpx_client): + """Verify correct URL, headers, JSON body, and MCPResponse parsing.""" + from openjarvis.mcp.transport import StreamableHTTPTransport + + mock_client = _mock_httpx_client + mock_client.post.return_value = _make_http_response({"tools": []}) + + transport = StreamableHTTPTransport("http://localhost:9583/mcp") + req = MCPRequest(method="tools/list", id=1) + resp = transport.send(req) + + # Verify the POST call + mock_client.post.assert_called_once() + call_kwargs = mock_client.post.call_args + assert call_kwargs[0][0] == "http://localhost:9583/mcp" + headers = call_kwargs[1]["headers"] + assert headers["Content-Type"] == "application/json" + assert "application/json" in headers["Accept"] + assert "text/event-stream" in headers["Accept"] + + # Verify JSON body matches the request + sent_json = call_kwargs[1]["json"] + assert sent_json["method"] == "tools/list" + assert sent_json["id"] == 1 + assert sent_json["jsonrpc"] == "2.0" + + # Verify response parsing + assert resp.error is None + assert resp.result == {"tools": []} + + def test_session_id_tracking(self, _mock_httpx_client): + """First response sets Mcp-Session-Id, subsequent requests include it.""" + from openjarvis.mcp.transport import StreamableHTTPTransport + + mock_client = _mock_httpx_client + # First response sets the session id + mock_client.post.return_value = _make_http_response( + {"capabilities": {}}, session_id="sess-abc-123" + ) + + transport = StreamableHTTPTransport("http://localhost:9583/mcp") + req1 = MCPRequest(method="initialize", id=1) + transport.send(req1) + + assert transport._session_id == "sess-abc-123" + + # Second request — prepare new response without session header + mock_client.post.return_value = _make_http_response({"tools": []}) + req2 = MCPRequest(method="tools/list", id=2) + transport.send(req2) + + # Verify the second call included the session id header + second_call_headers = mock_client.post.call_args[1]["headers"] + assert second_call_headers["Mcp-Session-Id"] == "sess-abc-123" + + def test_first_request_has_no_session_id(self, _mock_httpx_client): + """First request should not include Mcp-Session-Id header.""" + from openjarvis.mcp.transport import StreamableHTTPTransport + + mock_client = _mock_httpx_client + mock_client.post.return_value = _make_http_response({}) + + transport = StreamableHTTPTransport("http://localhost:9583/mcp") + transport.send(MCPRequest(method="initialize", id=1)) + + first_call_headers = mock_client.post.call_args[1]["headers"] + assert "Mcp-Session-Id" not in first_call_headers + + def test_connect_error_handling(self, _mock_httpx_client): + """httpx.ConnectError should be wrapped in RuntimeError.""" + import httpx + + from openjarvis.mcp.transport import StreamableHTTPTransport + + mock_client = _mock_httpx_client + mock_client.post.side_effect = httpx.ConnectError("Connection refused") + + transport = StreamableHTTPTransport("http://localhost:9583/mcp") + with pytest.raises(RuntimeError, match="Failed to connect"): + transport.send(MCPRequest(method="initialize", id=1)) + + def test_timeout_error_handling(self, _mock_httpx_client): + """httpx.TimeoutException should be wrapped in RuntimeError.""" + import httpx + + from openjarvis.mcp.transport import StreamableHTTPTransport + + mock_client = _mock_httpx_client + mock_client.post.side_effect = httpx.TimeoutException("Read timed out") + + transport = StreamableHTTPTransport("http://localhost:9583/mcp") + with pytest.raises(RuntimeError, match="Timeout communicating"): + transport.send(MCPRequest(method="initialize", id=1)) + + def test_close(self, _mock_httpx_client): + """close() should close the underlying httpx client.""" + from openjarvis.mcp.transport import StreamableHTTPTransport + + mock_client = _mock_httpx_client + transport = StreamableHTTPTransport("http://localhost:9583/mcp") + transport.close() + mock_client.close.assert_called_once() + + def test_backward_compat_alias(self): + """SSETransport should be the same class as StreamableHTTPTransport.""" + from openjarvis.mcp.transport import SSETransport, StreamableHTTPTransport + + assert SSETransport is StreamableHTTPTransport diff --git a/tests/mcp/test_transport.py b/tests/mcp/test_transport.py index f1cb545e..4cd10e86 100644 --- a/tests/mcp/test_transport.py +++ b/tests/mcp/test_transport.py @@ -5,13 +5,18 @@ from __future__ import annotations import json import sys import textwrap -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import pytest from openjarvis.mcp.protocol import MCPRequest from openjarvis.mcp.server import MCPServer -from openjarvis.mcp.transport import InProcessTransport, SSETransport, StdioTransport +from openjarvis.mcp.transport import ( + InProcessTransport, + SSETransport, + StdioTransport, + StreamableHTTPTransport, +) from openjarvis.tools.calculator import CalculatorTool from openjarvis.tools.think import ThinkTool @@ -156,65 +161,108 @@ class TestStdioTransport: transport.close() # Should not raise -class TestSSETransport: - def test_send_receive(self, monkeypatch): - """Mock httpx to simulate HTTP response.""" +class TestStreamableHTTPTransport: + """Tests for StreamableHTTPTransport (also aliased as SSETransport).""" + + def _make_mock_response(self, body: dict) -> MagicMock: + """Create a mock httpx.Response with the given JSON body.""" mock_response = MagicMock() - mock_response.text = json.dumps( + mock_response.text = json.dumps(body) + mock_response.headers = {} + mock_response.raise_for_status = MagicMock() + return mock_response + + @patch("httpx.Client") + def test_send_receive(self, mock_client_cls): + """Mock httpx.Client to simulate HTTP response.""" + mock_client = MagicMock() + mock_client_cls.return_value = mock_client + mock_client.post.return_value = self._make_mock_response( {"jsonrpc": "2.0", "id": 1, "result": {"tools": []}} ) - mock_response.raise_for_status = MagicMock() - mock_httpx = MagicMock() - mock_httpx.post.return_value = mock_response - monkeypatch.setitem(sys.modules, "httpx", mock_httpx) - - transport = SSETransport("http://localhost:8080/mcp") + transport = StreamableHTTPTransport("http://localhost:8080/mcp") req = MCPRequest(method="tools/list", id=1) resp = transport.send(req) assert resp.error is None assert resp.result == {"tools": []} - def test_send_posts_json(self, monkeypatch): + @patch("httpx.Client") + def test_send_posts_json(self, mock_client_cls): """Verify the HTTP POST includes correct headers and body.""" - mock_response = MagicMock() - mock_response.text = json.dumps({"jsonrpc": "2.0", "id": 1, "result": {}}) - mock_response.raise_for_status = MagicMock() + mock_client = MagicMock() + mock_client_cls.return_value = mock_client + mock_client.post.return_value = self._make_mock_response( + {"jsonrpc": "2.0", "id": 1, "result": {}} + ) - mock_httpx = MagicMock() - mock_httpx.post.return_value = mock_response - monkeypatch.setitem(sys.modules, "httpx", mock_httpx) - - transport = SSETransport("http://localhost:8080/mcp") + transport = StreamableHTTPTransport("http://localhost:8080/mcp") req = MCPRequest(method="initialize", id=1) transport.send(req) - call_args = mock_httpx.post.call_args + call_args = mock_client.post.call_args assert call_args[0][0] == "http://localhost:8080/mcp" assert call_args[1]["headers"]["Content-Type"] == "application/json" - def test_close_is_noop(self): - transport = SSETransport("http://localhost:8080/mcp") - transport.close() # Should not raise + @patch("httpx.Client") + def test_close_closes_client(self, mock_client_cls): + """close() should close the underlying httpx.Client.""" + mock_client = MagicMock() + mock_client_cls.return_value = mock_client - def test_error_response(self, monkeypatch): + transport = StreamableHTTPTransport("http://localhost:8080/mcp") + transport.close() + mock_client.close.assert_called_once() + + @patch("httpx.Client") + def test_error_response(self, mock_client_cls): """Simulate server returning an error response.""" - mock_response = MagicMock() - mock_response.text = json.dumps( + mock_client = MagicMock() + mock_client_cls.return_value = mock_client + mock_client.post.return_value = self._make_mock_response( { "jsonrpc": "2.0", "id": 1, "error": {"code": -32601, "message": "Not found"}, } ) - mock_response.raise_for_status = MagicMock() - mock_httpx = MagicMock() - mock_httpx.post.return_value = mock_response - monkeypatch.setitem(sys.modules, "httpx", mock_httpx) - - transport = SSETransport("http://localhost:8080/mcp") + transport = StreamableHTTPTransport("http://localhost:8080/mcp") req = MCPRequest(method="unknown", id=1) resp = transport.send(req) assert resp.error is not None assert resp.error["code"] == -32601 + + @patch("httpx.Client") + def test_session_id_tracking(self, mock_client_cls): + """Verify Mcp-Session-Id is tracked from response and sent on subsequent + requests.""" + mock_client = MagicMock() + mock_client_cls.return_value = mock_client + + # First response sets a session id + first_response = self._make_mock_response( + {"jsonrpc": "2.0", "id": 1, "result": {}} + ) + first_response.headers = {"mcp-session-id": "sess-abc-123"} + + # Second response + second_response = self._make_mock_response( + {"jsonrpc": "2.0", "id": 2, "result": {}} + ) + second_response.headers = {} + + mock_client.post.side_effect = [first_response, second_response] + + transport = StreamableHTTPTransport("http://localhost:8080/mcp") + transport.send(MCPRequest(method="initialize", id=1)) + assert transport._session_id == "sess-abc-123" + + transport.send(MCPRequest(method="tools/list", id=2)) + # Second call should include session id in headers + second_call_headers = mock_client.post.call_args_list[1][1]["headers"] + assert second_call_headers["Mcp-Session-Id"] == "sess-abc-123" + + def test_sse_transport_alias(self): + """SSETransport should be an alias for StreamableHTTPTransport.""" + assert SSETransport is StreamableHTTPTransport diff --git a/tests/server/test_agent_manager_routes.py b/tests/server/test_agent_manager_routes.py index b8dbd57c..dd908d67 100644 --- a/tests/server/test_agent_manager_routes.py +++ b/tests/server/test_agent_manager_routes.py @@ -244,39 +244,40 @@ def test_run_agent_concurrent_returns_409(tmp_path): @pytest.mark.skipif(not HAS_FASTAPI, reason="fastapi not installed") class TestAgentManagerStreaming: - """Tests for the SSE streaming mode of the managed-agent messages endpoint.""" + """Tests for the SSE streaming mode of the managed-agent messages endpoint. + + The new implementation uses engine.stream_full() for real token streaming + instead of agent.run() + word-by-word replay. + """ @pytest.fixture def _mock_engine(self): + """Create a mock engine with a working stream_full() method.""" + from openjarvis.engine._stubs import StreamChunk + engine = MagicMock() engine.engine_id = "mock" engine._model = "test-model" engine.health.return_value = True + + # Default stream_full: echo the last user message token-by-token + async def _stream_full(messages, *, model, **kwargs): + # Find the last user message content + last_content = "" + for m in reversed(messages): + if hasattr(m, "role") and m.role.value == "user": + last_content = m.content + break + response = f"Echo: {last_content}" + for token in response.split(" "): + yield StreamChunk(content=token + " ") + yield StreamChunk(finish_reason="stop") + + engine.stream_full = _stream_full return engine @pytest.fixture - def _mock_agent_cls(self): - """Register a mock agent class in the AgentRegistry for testing.""" - from openjarvis.agents._stubs import AgentResult - from openjarvis.core.registry import AgentRegistry - - class _MockStreamAgent: - agent_id = "mock_stream" - - def __init__(self, engine, model, **kwargs): - self._engine = engine - self._model = model - - def run(self, input_text, context=None, **kwargs): - return AgentResult(content=f"Echo: {input_text}", turns=1) - - # Register under a unique key for test isolation - AgentRegistry._entries()["_test_stream"] = _MockStreamAgent - yield _MockStreamAgent - AgentRegistry._entries().pop("_test_stream", None) - - @pytest.fixture - def stream_client(self, manager, _mock_engine, _mock_agent_cls): + def stream_client(self, manager, _mock_engine): from fastapi import FastAPI from openjarvis.server.agent_manager_routes import create_agent_manager_router @@ -293,10 +294,11 @@ class TestAgentManagerStreaming: app.include_router(tools_router) return TestClient(app) - def test_send_message_stream(self, manager, stream_client, _mock_agent_cls): + def test_send_message_stream(self, manager, stream_client): """Test streaming mode returns SSE response with [DONE] sentinel.""" agent = manager.create_agent( - name="streamer", agent_type="_test_stream", + name="streamer", + agent_type="simple", ) resp = stream_client.post( f"/v1/managed-agents/{agent['id']}/messages", @@ -312,10 +314,11 @@ class TestAgentManagerStreaming: # Last data line must be [DONE] assert data_lines[-1].strip() == "data: [DONE]" - def test_send_message_stream_content(self, manager, stream_client, _mock_agent_cls): - """Test streaming returns the correct agent response content.""" + def test_send_message_stream_real_tokens(self, manager, stream_client): + """Content arrives as real tokens, not word-burst after completion.""" agent = manager.create_agent( - name="streamer2", agent_type="_test_stream", + name="streamer_tokens", + agent_type="simple", ) resp = stream_client.post( f"/v1/managed-agents/{agent['id']}/messages", @@ -324,7 +327,7 @@ class TestAgentManagerStreaming: assert resp.status_code == 200 # Collect content tokens from stream - content = "" + content_chunks = [] for line in resp.text.strip().split("\n"): if line.startswith("data:") and "[DONE]" not in line: raw = line[5:].strip() @@ -335,16 +338,19 @@ class TestAgentManagerStreaming: choices = data.get("choices", [{}]) delta_content = choices[0].get("delta", {}).get("content") if delta_content: - content += delta_content + content_chunks.append(delta_content) - assert content == "Echo: Hello world" + # Should have multiple token chunks (real streaming, not single burst) + assert len(content_chunks) > 1 + full_content = "".join(content_chunks) + assert "Echo:" in full_content + assert "Hello world" in full_content - def test_send_message_stream_stores_response( - self, manager, stream_client, _mock_agent_cls, - ): + def test_send_message_stream_stores_response(self, manager, stream_client): """After streaming, agent response is persisted in the DB.""" agent = manager.create_agent( - name="streamer3", agent_type="_test_stream", + name="streamer3", + agent_type="simple", ) resp = stream_client.post( f"/v1/managed-agents/{agent['id']}/messages", @@ -362,12 +368,11 @@ class TestAgentManagerStreaming: agent_msg = next(m for m in messages if m["direction"] == "agent_to_user") assert "persist me" in agent_msg["content"] - def test_send_message_stream_finish_reason( - self, manager, stream_client, _mock_agent_cls, - ): + def test_send_message_stream_finish_reason(self, manager, stream_client): """The final chunk before [DONE] has finish_reason='stop'.""" agent = manager.create_agent( - name="streamer4", agent_type="_test_stream", + name="streamer4", + agent_type="simple", ) resp = stream_client.post( f"/v1/managed-agents/{agent['id']}/messages", @@ -385,3 +390,38 @@ class TestAgentManagerStreaming: # Last chunk should have finish_reason="stop" assert chunks[-1]["choices"][0]["finish_reason"] == "stop" + + def test_send_message_stream_error_handling(self, manager): + """Engine errors are reported gracefully via SSE.""" + + error_engine = MagicMock() + error_engine.engine_id = "error" + error_engine._model = "test-model" + + async def _stream_full_error(messages, *, model, **kwargs): + raise RuntimeError("LLM connection failed") + yield # make it a generator # noqa: E501 + + error_engine.stream_full = _stream_full_error + + from fastapi import FastAPI + from fastapi.testclient import TestClient as TC + + from openjarvis.server.agent_manager_routes import create_agent_manager_router + + app = FastAPI() + app.state.engine = error_engine + app.state.bus = None + routers = create_agent_manager_router(manager) + for r in routers: + app.include_router(r) + client = TC(app) + + agent = manager.create_agent(name="err_agent", agent_type="simple") + resp = client.post( + f"/v1/managed-agents/{agent['id']}/messages", + json={"content": "fail", "stream": True}, + ) + assert resp.status_code == 200 + assert "Error:" in resp.text or "error" in resp.text.lower() + assert "data: [DONE]" in resp.text diff --git a/tests/server/test_mcp_tools_cache.py b/tests/server/test_mcp_tools_cache.py new file mode 100644 index 00000000..d3bcdfa4 --- /dev/null +++ b/tests/server/test_mcp_tools_cache.py @@ -0,0 +1,154 @@ +"""Tests for _get_mcp_tools() caching in agent_manager_routes.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import pytest + +pytest.importorskip("fastapi", reason="fastapi required for server route tests") + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +class _FakeAppState: + """Minimal app_state substitute with dynamic attributes.""" + + pass + + +def _make_config(*, enabled: bool = True, servers_json: str = "[]") -> MagicMock: + """Build a mock config with tools.mcp.enabled and tools.mcp.servers.""" + config = MagicMock() + config.tools.mcp.enabled = enabled + config.tools.mcp.servers = servers_json + return config + + +def _make_tool_spec(name: str, description: str = "") -> MagicMock: + spec = MagicMock() + spec.name = name + spec.description = description + spec.parameters = {"type": "object", "properties": {}} + return spec + + +def _make_adapter(name: str) -> MagicMock: + adapter = MagicMock() + adapter.spec = _make_tool_spec(name) + return adapter + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +@patch("openjarvis.core.config.load_config") +def test_returns_tools_from_mcp_server(mock_load_config: MagicMock): + """With a mocked MCP server, discovered tools are returned.""" + from openjarvis.server.agent_manager_routes import _get_mcp_tools + + server_cfg = [{"name": "test-server", "url": "http://localhost:9999"}] + mock_load_config.return_value = _make_config( + servers_json=json.dumps(server_cfg), + ) + + mock_adapter = _make_adapter("get_weather") + + with ( + patch("openjarvis.mcp.transport.StreamableHTTPTransport"), + patch("openjarvis.mcp.client.MCPClient"), + patch("openjarvis.tools.mcp_adapter.MCPToolProvider") as MockProvider, + ): + MockProvider.return_value.discover.return_value = [mock_adapter] + + app_state = _FakeAppState() + tools, adapters = _get_mcp_tools(app_state) + + assert len(tools) == 1 + assert tools[0]["function"]["name"] == "get_weather" + assert "get_weather" in adapters + + +@patch("openjarvis.core.config.load_config") +def test_caches_successful_discovery(mock_load_config: MagicMock): + """Second call returns cached result without re-discovering.""" + from openjarvis.server.agent_manager_routes import _get_mcp_tools + + server_cfg = [{"name": "test-server", "url": "http://localhost:9999"}] + mock_load_config.return_value = _make_config( + servers_json=json.dumps(server_cfg), + ) + + mock_adapter = _make_adapter("cached_tool") + + with ( + patch("openjarvis.mcp.transport.StreamableHTTPTransport"), + patch("openjarvis.mcp.client.MCPClient"), + patch("openjarvis.tools.mcp_adapter.MCPToolProvider") as MockProvider, + ): + MockProvider.return_value.discover.return_value = [mock_adapter] + + app_state = _FakeAppState() + + # First call discovers + tools1, _ = _get_mcp_tools(app_state) + assert len(tools1) == 1 + + # Second call should use cache (discover not called again) + discover_call_count = MockProvider.return_value.discover.call_count + tools2, _ = _get_mcp_tools(app_state) + assert len(tools2) == 1 + assert MockProvider.return_value.discover.call_count == discover_call_count + + +@patch("openjarvis.core.config.load_config") +def test_does_not_cache_empty_results(mock_load_config: MagicMock): + """Failed/empty discovery is not cached so it can be retried.""" + from openjarvis.server.agent_manager_routes import _get_mcp_tools + + server_cfg = [{"name": "failing-server", "url": "http://localhost:9999"}] + mock_load_config.return_value = _make_config( + servers_json=json.dumps(server_cfg), + ) + + with ( + patch("openjarvis.mcp.transport.StreamableHTTPTransport"), + patch("openjarvis.mcp.client.MCPClient"), + patch("openjarvis.tools.mcp_adapter.MCPToolProvider") as MockProvider, + ): + # First call: discovery returns empty + MockProvider.return_value.discover.return_value = [] + app_state = _FakeAppState() + + tools1, _ = _get_mcp_tools(app_state) + assert len(tools1) == 0 + + # Verify no cache was set (empty result) + assert getattr(app_state, "_mcp_tools_cache", None) is None + + # Second call: discovery now returns something + mock_adapter = _make_adapter("retry_tool") + MockProvider.return_value.discover.return_value = [mock_adapter] + + tools2, _ = _get_mcp_tools(app_state) + assert len(tools2) == 1 + assert tools2[0]["function"]["name"] == "retry_tool" + + +@patch("openjarvis.core.config.load_config") +def test_handles_config_load_failure(mock_load_config: MagicMock): + """Config load failure returns empty, no crash.""" + from openjarvis.server.agent_manager_routes import _get_mcp_tools + + mock_load_config.side_effect = RuntimeError("config broken") + + app_state = _FakeAppState() + tools, adapters = _get_mcp_tools(app_state) + + assert tools == [] + assert adapters == {}