From 76c6086db3eb73f406e0f0163647532ecade3397 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sat, 25 Jul 2026 12:50:39 +0000 Subject: [PATCH] feat: MCP/OpenAPI tool calling support --- .env.example | 6 + steward/bot/telegram.py | 45 ++-- steward/config.py | 6 + steward/llm/client.py | 104 ++++++++ steward/main.py | 17 +- steward/tools/__init__.py | 1 + steward/tools/client.py | 259 ++++++++++++++++++++ tests/test_tools.py | 499 ++++++++++++++++++++++++++++++++++++++ 8 files changed, 921 insertions(+), 16 deletions(-) create mode 100644 steward/tools/__init__.py create mode 100644 steward/tools/client.py create mode 100644 tests/test_tools.py diff --git a/.env.example b/.env.example index 4e4571f..d214b4d 100644 --- a/.env.example +++ b/.env.example @@ -24,3 +24,9 @@ OPENAI_MODEL=gpt-4o # Thread memory store path (JSON file for persisted thread summaries) # THREAD_MEMORY_PATH=thread_memory.json + +# MCP / OpenAPI tool server (open-webui/openapi-servers compatible) +# Point to any OpenAPI-spec tool server to enable LLM tool calling. +# The service fetches /openapi.json from MCP_SERVER_URL to discover tools. +# MCP_SERVER_URL=http://localhost:8000 +# MCP_SERVER_API_KEY= diff --git a/steward/bot/telegram.py b/steward/bot/telegram.py index 8c1e072..a56b47a 100644 --- a/steward/bot/telegram.py +++ b/steward/bot/telegram.py @@ -2,6 +2,7 @@ import logging from collections import defaultdict +from typing import Any from telegram import Update from telegram.constants import ParseMode @@ -17,17 +18,18 @@ from steward.config import Settings from steward.llm.client import LLMClient from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary from steward.proposals.generator import Proposal, ProposalGenerator +from steward.tools.client import ToolClient logger = logging.getLogger(__name__) # Per-user conversation history for non-threaded messages (capped at _MAX_HISTORY turns). -_history: dict[int, list[dict[str, str]]] = defaultdict(list) +_history: dict[int, list[dict[str, Any]]] = defaultdict(list) _MAX_HISTORY = 20 # Per-thread conversation history: (chat_id, thread_id) → full history (unbounded). # Messages belonging to a Telegram message thread are kept in their entirety here # until explicitly flushed by the /flush command. -_thread_history: dict[tuple[int, int], list[dict[str, str]]] = defaultdict(list) +_thread_history: dict[tuple[int, int], list[dict[str, Any]]] = defaultdict(list) _FLUSH_SYSTEM_PROMPT = ( "You are Steward. The following is a complete Telegram message thread conversation. " @@ -318,9 +320,9 @@ async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> def _with_kb_context( - history: list[dict[str, str]], + history: list[dict[str, Any]], relevant: list[ThreadSummary], -) -> list[dict[str, str]]: +) -> list[dict[str, Any]]: """Prepend relevant knowledge-base summaries as a transient system context message. The returned list is a *new* list — the original *history* is not mutated. @@ -344,10 +346,15 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> overlap with keywords in the current message. Any matches are injected as transient context — they are NOT stored in the rolling history, so they do not permanently consume the context window. + + If a :class:`~steward.tools.client.ToolClient` is registered in + ``context.bot_data``, the LLM is invoked with tool calling support so it + can take actions on the configured MCP/OpenAPI tool server. """ settings: Settings = context.bot_data["settings"] llm: LLMClient = context.bot_data["llm"] store: ThreadMemoryStore = context.bot_data["thread_store"] + tool_client: ToolClient | None = context.bot_data.get("tool_client") user = update.effective_user if user is None or not _is_allowed(user.id, settings): return @@ -359,32 +366,42 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> key = _thread_key(update) if key is not None: # Thread message: unbounded history - history = _thread_history[key] + history: list[dict[str, Any]] = _thread_history[key] call_history = _with_kb_context(history, store.search(text)) - reply = await llm.chat(text, history=call_history) + if tool_client is not None: + reply = await llm.chat_with_tools(text, tool_client, history=call_history) + else: + reply = await llm.chat(text, history=call_history) # type: ignore[arg-type] history.append({"role": "user", "content": text}) history.append({"role": "assistant", "content": reply}) else: # Non-thread message: capped history per user - history = _history[user.id] - call_history = _with_kb_context(history, store.search(text)) - reply = await llm.chat(text, history=call_history) - history.append({"role": "user", "content": text}) - history.append({"role": "assistant", "content": reply}) - if len(history) > _MAX_HISTORY * 2: - _history[user.id] = history[-(_MAX_HISTORY * 2) :] + user_history: list[dict[str, Any]] = _history[user.id] + call_history = _with_kb_context(user_history, store.search(text)) + if tool_client is not None: + reply = await llm.chat_with_tools(text, tool_client, history=call_history) + else: + reply = await llm.chat(text, history=call_history) # type: ignore[arg-type] + user_history.append({"role": "user", "content": text}) + user_history.append({"role": "assistant", "content": reply}) + if len(user_history) > _MAX_HISTORY * 2: + _history[user.id] = user_history[-(_MAX_HISTORY * 2) :] await _send_long(update, reply) def build_application( - settings: Settings, llm: LLMClient, thread_store: ThreadMemoryStore | None = None + settings: Settings, + llm: LLMClient, + thread_store: ThreadMemoryStore | None = None, + tool_client: ToolClient | None = None, ) -> Application: # type: ignore[type-arg] """Build and return the Telegram Application.""" app = Application.builder().token(settings.telegram_bot_token).build() app.bot_data["settings"] = settings app.bot_data["llm"] = llm app.bot_data["thread_store"] = thread_store or ThreadMemoryStore(settings.thread_memory_path) + app.bot_data["tool_client"] = tool_client # None when tools are not configured app.add_handler(CommandHandler("start", start_handler)) app.add_handler(CommandHandler("help", help_handler)) diff --git a/steward/config.py b/steward/config.py index a1a70d5..aac14b5 100644 --- a/steward/config.py +++ b/steward/config.py @@ -32,6 +32,12 @@ class Settings(BaseSettings): # Thread memory thread_memory_path: str = "thread_memory.json" + # MCP / OpenAPI tool server (open-webui/openapi-servers compatible) + # Set MCP_SERVER_URL to enable tool calling. The service fetches + # /openapi.json from this URL to discover available tools. + mcp_server_url: str = "" + mcp_server_api_key: str = "" + def get_settings() -> Settings: """Return application settings singleton.""" diff --git a/steward/llm/client.py b/steward/llm/client.py index 883cd4a..a048d2d 100644 --- a/steward/llm/client.py +++ b/steward/llm/client.py @@ -1,12 +1,17 @@ """LLM client wrapper (OpenAI-compatible).""" +import json import logging from collections.abc import AsyncIterator +from typing import TYPE_CHECKING, Any from openai import AsyncOpenAI from steward.config import Settings +if TYPE_CHECKING: + from steward.tools.client import ToolClient + logger = logging.getLogger(__name__) @@ -45,6 +50,105 @@ class LLMClient: logger.debug("LLM reply: %d chars", len(reply)) return reply + async def chat_with_tools( + self, + user_message: str, + tool_client: "ToolClient", + *, + history: list[dict[str, Any]] | None = None, + system_prompt: str | None = None, + max_tool_rounds: int = 10, + ) -> str: + """Chat with tool calling support (agentic loop). + + Sends *user_message* to the LLM with the tool definitions supplied by + *tool_client*. If the model requests one or more tool calls, each call + is executed via *tool_client* and the results are fed back to the model. + The loop repeats until the model returns a plain-text response or + *max_tool_rounds* is reached. + + Intermediate tool-call messages are **not** returned to the caller and + should **not** be stored in the persistent conversation history; only + the final assistant reply needs to be recorded alongside the original + user message. + """ + tools = await tool_client.get_tools() + system = system_prompt or self._settings.openai_system_prompt + messages: list[dict[str, Any]] = [{"role": "system", "content": system}] + if history: + messages.extend(history) + messages.append({"role": "user", "content": user_message}) + + logger.debug( + "LLM tool-call request: model=%s tools=%d messages=%d", + self._settings.openai_model, + len(tools), + len(messages), + ) + + last_content = "" + for round_num in range(max_tool_rounds): + kwargs: dict[str, Any] = { + "model": self._settings.openai_model, + "messages": messages, + } + if tools: + kwargs["tools"] = tools + kwargs["tool_choice"] = "auto" + + response = await self._client.chat.completions.create(**kwargs) # type: ignore[arg-type] + choice = response.choices[0] + msg = choice.message + last_content = msg.content or "" + + if not msg.tool_calls: + # No tool calls → final answer + logger.debug("LLM reply (round %d): %d chars", round_num, len(last_content)) + return last_content + + tool_names = [tc.function.name for tc in msg.tool_calls] + logger.info("Tool calls in round %d: %s", round_num + 1, tool_names) + + # Append the assistant's tool-call turn + messages.append( + { + "role": "assistant", + "content": msg.content, + "tool_calls": [ + { + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments, + }, + } + for tc in msg.tool_calls + ], + } + ) + + # Execute each tool and append results + for tc in msg.tool_calls: + try: + args = json.loads(tc.function.arguments) + result = await tool_client.call(tc.function.name, args) + logger.info("Tool %s → %d chars", tc.function.name, len(result)) + except Exception as exc: + result = f"Tool call failed: {exc}" + logger.warning("Tool %s failed: %s", tc.function.name, exc) + + messages.append( + { + "role": "tool", + "tool_call_id": tc.id, + "content": result, + } + ) + + logger.warning("Reached max tool rounds (%d); returning last content", max_tool_rounds) + return last_content + async def stream( self, user_message: str, diff --git a/steward/main.py b/steward/main.py index 8d0cb9c..02845ae 100644 --- a/steward/main.py +++ b/steward/main.py @@ -10,6 +10,7 @@ from steward.config import get_settings from steward.llm.client import LLMClient from steward.memory.thread_store import ThreadMemoryStore from steward.proposals.generator import ProposalGenerator +from steward.tools.client import ToolClient logging.basicConfig( level=logging.INFO, @@ -46,7 +47,18 @@ def main() -> None: llm = LLMClient(settings) thread_store = ThreadMemoryStore(settings.thread_memory_path) - app = build_application(settings, llm, thread_store) + + tool_client: ToolClient | None = None + if settings.mcp_server_url: + tool_client = ToolClient( + base_url=settings.mcp_server_url, + api_key=settings.mcp_server_api_key, + ) + logger.info("MCP tool server configured: %s", settings.mcp_server_url) + else: + logger.info("No MCP_SERVER_URL configured – tool calling disabled") + + app = build_application(settings, llm, thread_store, tool_client) generator = ProposalGenerator(settings, llm) scheduler = AsyncIOScheduler() @@ -60,9 +72,10 @@ def main() -> None: scheduler.start() logger.info( - "Steward starting: model=%s analysis_url=%s", + "Steward starting: model=%s analysis_url=%s tools=%s", settings.openai_model, settings.analysis_target_url or "(none)", + settings.mcp_server_url or "(none)", ) app.run_polling(allowed_updates=["message"]) diff --git a/steward/tools/__init__.py b/steward/tools/__init__.py new file mode 100644 index 0000000..679d391 --- /dev/null +++ b/steward/tools/__init__.py @@ -0,0 +1 @@ +"""MCP-compatible OpenAPI tool calling support.""" diff --git a/steward/tools/client.py b/steward/tools/client.py new file mode 100644 index 0000000..062a2d8 --- /dev/null +++ b/steward/tools/client.py @@ -0,0 +1,259 @@ +"""OpenAPI-based tool client for MCP-compatible tool servers. + +Compatible with open-webui/openapi-servers (https://github.com/open-webui/openapi-servers). +Fetches the server's OpenAPI spec, converts operations to OpenAI function-calling +definitions, and executes tool calls by making HTTP requests back to the server. +""" + +from __future__ import annotations + +import json +import logging +from typing import Any + +import httpx + +logger = logging.getLogger(__name__) + + +class ToolClient: + """Client for OpenAPI-compatible tool servers. + + Usage:: + + client = ToolClient(base_url="http://localhost:8000", api_key="…") + tools = await client.get_tools() # OpenAI tool definitions + result = await client.call("my_op", {}) # execute a tool call + """ + + def __init__(self, base_url: str, api_key: str = "", timeout: float = 30.0) -> None: + self._base_url = base_url.rstrip("/") + self._api_key = api_key + self._timeout = timeout + self._spec: dict[str, Any] | None = None + self._tools: list[dict[str, Any]] | None = None + + @property + def base_url(self) -> str: + return self._base_url + + def _headers(self) -> dict[str, str]: + headers: dict[str, str] = {"Accept": "application/json"} + if self._api_key: + headers["Authorization"] = "Bearer " + self._api_key + return headers + + async def get_spec(self) -> dict[str, Any]: + """Fetch (and cache) the OpenAPI spec from ``/openapi.json``.""" + if self._spec is not None: + return self._spec + async with httpx.AsyncClient(timeout=self._timeout) as http: + resp = await http.get( + f"{self._base_url}/openapi.json", + headers=self._headers(), + ) + resp.raise_for_status() + self._spec = resp.json() + path_count = len(self._spec.get("paths", {})) + logger.info("Loaded OpenAPI spec from %s (%d paths)", self._base_url, path_count) + return self._spec + + async def get_tools(self) -> list[dict[str, Any]]: + """Return OpenAI function-calling tool definitions from the OpenAPI spec.""" + if self._tools is not None: + return self._tools + spec = await self.get_spec() + self._tools = _spec_to_openai_tools(spec) + logger.info("Registered %d tools from %s", len(self._tools), self._base_url) + return self._tools + + async def call(self, tool_name: str, arguments: dict[str, Any]) -> str: + """Execute a named tool call and return the result as a string. + + Resolves *tool_name* to a path+method in the cached OpenAPI spec, + partitions *arguments* into path params, query params, and request + body, then makes the HTTP request. + """ + spec = await self.get_spec() + path, method, path_params, query_params, body = _resolve_operation( + spec, tool_name, arguments + ) + + url = self._base_url + path + for key, value in path_params.items(): + url = url.replace(f"{{{key}}}", str(value)) + + async with httpx.AsyncClient(timeout=self._timeout) as http: + resp = await http.request( + method=method.upper(), + url=url, + headers=self._headers(), + params=query_params if query_params else None, + json=body if body else None, + ) + resp.raise_for_status() + content_type = resp.headers.get("content-type", "") + if "application/json" in content_type: + try: + return json.dumps(resp.json(), indent=2) + except ValueError: + pass + return resp.text + + +# --------------------------------------------------------------------------- +# OpenAPI → OpenAI tool-definition helpers +# --------------------------------------------------------------------------- + + +def _resolve_ref( + schema: dict[str, Any], components: dict[str, Any] +) -> dict[str, Any]: + """Recursively resolve a ``$ref`` inside an OpenAPI schema.""" + if "$ref" not in schema: + return schema + ref: str = schema["$ref"] # e.g. "#/components/schemas/Foo" + parts = ref.lstrip("#/").split("/") + # parts = ["components", "schemas", "Foo"] + obj: Any = {"components": components} + for part in parts: + if not isinstance(obj, dict): + return {} + obj = obj.get(part, {}) + return _schema_to_json_schema(obj, components) + + +def _schema_to_json_schema( + schema: dict[str, Any], components: dict[str, Any] | None = None +) -> dict[str, Any]: + """Convert an OpenAPI schema object to a JSON Schema-compatible dict.""" + if not schema: + return {} + resolved = _resolve_ref(schema, components or {}) if "$ref" in schema else schema + result: dict[str, Any] = {} + for key in ("type", "description", "enum", "format", "default", "minimum", "maximum"): + if key in resolved: + result[key] = resolved[key] + if "items" in resolved: + result["items"] = _schema_to_json_schema(resolved["items"], components) + if "properties" in resolved: + result["properties"] = { + k: _schema_to_json_schema(v, components) + for k, v in resolved["properties"].items() + } + if "required" in resolved: + result["required"] = resolved["required"] + return result + + +def _operation_id(method: str, path: str, op: dict[str, Any]) -> str: + """Return an operation ID, generating one from method+path if absent.""" + if op.get("operationId"): + return str(op["operationId"]) + slug = path.strip("/").replace("/", "_").replace("{", "").replace("}", "") + return f"{method}_{slug}" if slug else method + + +def _spec_to_openai_tools(spec: dict[str, Any]) -> list[dict[str, Any]]: + """Convert an OpenAPI 3.x spec to a list of OpenAI function-calling tool defs.""" + tools: list[dict[str, Any]] = [] + paths: dict[str, Any] = spec.get("paths", {}) + components: dict[str, Any] = spec.get("components", {}) + + for path, path_item in paths.items(): + for http_method in ("get", "post", "put", "patch", "delete"): + op: dict[str, Any] | None = path_item.get(http_method) + if op is None: + continue + + op_id = _operation_id(http_method, path, op) + description: str = op.get("summary") or op.get("description") or "" + + properties: dict[str, Any] = {} + required: list[str] = [] + + # URL / query parameters + for param in op.get("parameters", []): + name: str = param["name"] + schema = _schema_to_json_schema( + param.get("schema", {"type": "string"}), components + ) + if param.get("description"): + schema["description"] = param["description"] + properties[name] = schema + if param.get("required"): + required.append(name) + + # Request body – flatten top-level object properties + rb: dict[str, Any] = op.get("requestBody", {}) + if rb: + json_content = rb.get("content", {}).get("application/json", {}) + body_schema = _schema_to_json_schema( + json_content.get("schema", {}), components + ) + if body_schema.get("type") == "object": + for prop_name, prop_schema in body_schema.get("properties", {}).items(): + properties[prop_name] = prop_schema + if rb.get("required", False): + required.extend(body_schema.get("required", [])) + + tool: dict[str, Any] = { + "type": "function", + "function": { + "name": op_id, + "description": description, + "parameters": { + "type": "object", + "properties": properties, + }, + }, + } + if required: + tool["function"]["parameters"]["required"] = required + + tools.append(tool) + + return tools + + +def _resolve_operation( + spec: dict[str, Any], + operation_id: str, + arguments: dict[str, Any], +) -> tuple[str, str, dict[str, Any], dict[str, Any], dict[str, Any]]: + """Find path+method for *operation_id* and split *arguments* into parts. + + Returns ``(path, method, path_params, query_params, body)``. + + * ``path_params`` – values that replace ``{placeholders}`` in the URL path. + * ``query_params`` – values appended as query string. + * ``body`` – any remaining arguments sent as the JSON request body. + """ + paths: dict[str, Any] = spec.get("paths", {}) + + for path, path_item in paths.items(): + for http_method in ("get", "post", "put", "patch", "delete"): + op: dict[str, Any] | None = path_item.get(http_method) + if op is None: + continue + if _operation_id(http_method, path, op) != operation_id: + continue + + path_params: dict[str, Any] = {} + query_params: dict[str, Any] = {} + declared_params: set[str] = set() + + for param in op.get("parameters", []): + name = param["name"] + declared_params.add(name) + if name not in arguments: + continue + if param.get("in") == "path": + path_params[name] = arguments[name] + else: + query_params[name] = arguments[name] + + body = {k: v for k, v in arguments.items() if k not in declared_params} + return path, http_method, path_params, query_params, body + + raise ValueError(f"Operation '{operation_id}' not found in spec") diff --git a/tests/test_tools.py b/tests/test_tools.py new file mode 100644 index 0000000..d7e01a7 --- /dev/null +++ b/tests/test_tools.py @@ -0,0 +1,499 @@ +"""Tests for steward.tools.client and LLM tool-calling integration.""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +import respx +from httpx import Response + +from steward.config import Settings +from steward.llm.client import LLMClient +from steward.tools.client import ( + ToolClient, + _operation_id, + _resolve_operation, + _schema_to_json_schema, + _spec_to_openai_tools, +) + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +SIMPLE_SPEC: dict[str, Any] = { + "openapi": "3.0.0", + "info": {"title": "Test", "version": "1.0.0"}, + "paths": { + "/items": { + "get": { + "operationId": "list_items", + "summary": "List all items", + "parameters": [ + { + "name": "limit", + "in": "query", + "required": False, + "schema": {"type": "integer"}, + "description": "Max results", + } + ], + "responses": {"200": {"description": "OK"}}, + } + }, + "/items/{item_id}": { + "get": { + "operationId": "get_item", + "summary": "Get a single item", + "parameters": [ + { + "name": "item_id", + "in": "path", + "required": True, + "schema": {"type": "string"}, + } + ], + "responses": {"200": {"description": "OK"}}, + } + }, + "/items/create": { + "post": { + "operationId": "create_item", + "summary": "Create an item", + "requestBody": { + "required": True, + "content": { + "application/json": { + "schema": { + "type": "object", + "properties": { + "name": {"type": "string", "description": "Item name"}, + "value": {"type": "integer"}, + }, + "required": ["name"], + } + } + }, + }, + "responses": {"201": {"description": "Created"}}, + } + }, + }, +} + + +def _make_llm_settings(**kwargs: Any) -> Settings: + defaults = dict( + telegram_bot_token="t", + openai_api_key="k", + openai_model="gpt-4o", + openai_system_prompt="You are Steward.", + ) + defaults.update(kwargs) + return Settings(**defaults) + + +# --------------------------------------------------------------------------- +# _operation_id +# --------------------------------------------------------------------------- + + +def test_generate_operation_id_simple(): + op: dict[str, Any] = {} + assert _operation_id("get", "/items", op) == "get_items" + + +def test_generate_operation_id_with_path_param(): + op: dict[str, Any] = {} + assert _operation_id("get", "/items/{item_id}", op) == "get_items_item_id" + + +def test_generate_operation_id_root(): + op: dict[str, Any] = {} + assert _operation_id("get", "/", op) == "get" + + +def test_operation_id_uses_declared(): + op = {"operationId": "my_op"} + assert _operation_id("get", "/foo", op) == "my_op" + + +def test_operation_id_generates_from_path(): + op: dict[str, Any] = {} + assert _operation_id("post", "/foo/bar", op) == "post_foo_bar" + + +# --------------------------------------------------------------------------- +# _schema_to_json_schema +# --------------------------------------------------------------------------- + + +def test_schema_to_json_schema_basic(): + schema = {"type": "string", "description": "A name"} + result = _schema_to_json_schema(schema) + assert result == {"type": "string", "description": "A name"} + + +def test_schema_to_json_schema_resolves_ref(): + components = {"schemas": {"Foo": {"type": "integer"}}} + schema = {"$ref": "#/components/schemas/Foo"} + result = _schema_to_json_schema(schema, components) + assert result == {"type": "integer"} + + +def test_schema_to_json_schema_nested_properties(): + schema = { + "type": "object", + "properties": { + "name": {"type": "string"}, + "count": {"type": "integer"}, + }, + "required": ["name"], + } + result = _schema_to_json_schema(schema) + assert result["type"] == "object" + assert "name" in result["properties"] + assert result["required"] == ["name"] + + +# --------------------------------------------------------------------------- +# _spec_to_openai_tools +# --------------------------------------------------------------------------- + + +def test_spec_to_openai_tools_count(): + tools = _spec_to_openai_tools(SIMPLE_SPEC) + assert len(tools) == 3 + + +def test_spec_to_openai_tools_structure(): + tools = _spec_to_openai_tools(SIMPLE_SPEC) + tool = next(t for t in tools if t["function"]["name"] == "list_items") + assert tool["type"] == "function" + assert tool["function"]["description"] == "List all items" + params = tool["function"]["parameters"] + assert params["type"] == "object" + assert "limit" in params["properties"] + + +def test_spec_to_openai_tools_path_param_required(): + tools = _spec_to_openai_tools(SIMPLE_SPEC) + tool = next(t for t in tools if t["function"]["name"] == "get_item") + assert "item_id" in tool["function"]["parameters"]["properties"] + assert "item_id" in tool["function"]["parameters"]["required"] + + +def test_spec_to_openai_tools_request_body(): + tools = _spec_to_openai_tools(SIMPLE_SPEC) + tool = next(t for t in tools if t["function"]["name"] == "create_item") + props = tool["function"]["parameters"]["properties"] + assert "name" in props + assert "value" in props + + +def test_spec_to_openai_tools_empty_spec(): + assert _spec_to_openai_tools({}) == [] + + +# --------------------------------------------------------------------------- +# _resolve_operation +# --------------------------------------------------------------------------- + + +def test_resolve_operation_query_params(): + path, method, path_p, query_p, body = _resolve_operation( + SIMPLE_SPEC, "list_items", {"limit": 5} + ) + assert path == "/items" + assert method == "get" + assert path_p == {} + assert query_p == {"limit": 5} + assert body == {} + + +def test_resolve_operation_path_params(): + path, method, path_p, query_p, body = _resolve_operation( + SIMPLE_SPEC, "get_item", {"item_id": "abc"} + ) + assert path == "/items/{item_id}" + assert method == "get" + assert path_p == {"item_id": "abc"} + assert query_p == {} + assert body == {} + + +def test_resolve_operation_body(): + path, method, path_p, query_p, body = _resolve_operation( + SIMPLE_SPEC, "create_item", {"name": "foo", "value": 42} + ) + assert path == "/items/create" + assert method == "post" + assert path_p == {} + assert query_p == {} + assert body == {"name": "foo", "value": 42} + + +def test_resolve_operation_not_found(): + with pytest.raises(ValueError, match="not found"): + _resolve_operation(SIMPLE_SPEC, "nonexistent_op", {}) + + +# --------------------------------------------------------------------------- +# ToolClient +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_tool_client_get_spec(): + """get_spec() fetches from /openapi.json and caches the result.""" + with respx.mock: + respx.get("http://tools.local/openapi.json").mock( + return_value=Response(200, json=SIMPLE_SPEC) + ) + client = ToolClient("http://tools.local") + spec = await client.get_spec() + assert spec["openapi"] == "3.0.0" + # Second call uses cache – no additional HTTP request + spec2 = await client.get_spec() + assert spec2 is spec + + +@pytest.mark.asyncio +async def test_tool_client_get_tools(): + with respx.mock: + respx.get("http://tools.local/openapi.json").mock( + return_value=Response(200, json=SIMPLE_SPEC) + ) + client = ToolClient("http://tools.local") + tools = await client.get_tools() + assert len(tools) == 3 + # Cached + tools2 = await client.get_tools() + assert tools2 is tools + + +@pytest.mark.asyncio +async def test_tool_client_call_get_with_query(): + with respx.mock: + respx.get("http://tools.local/openapi.json").mock( + return_value=Response(200, json=SIMPLE_SPEC) + ) + respx.get("http://tools.local/items").mock( + return_value=Response(200, json=[{"id": 1}]) + ) + client = ToolClient("http://tools.local") + result = await client.call("list_items", {"limit": 10}) + assert "id" in result + + +@pytest.mark.asyncio +async def test_tool_client_call_path_substitution(): + with respx.mock: + respx.get("http://tools.local/openapi.json").mock( + return_value=Response(200, json=SIMPLE_SPEC) + ) + respx.get("http://tools.local/items/abc123").mock( + return_value=Response(200, json={"id": "abc123"}) + ) + client = ToolClient("http://tools.local") + result = await client.call("get_item", {"item_id": "abc123"}) + assert "abc123" in result + + +@pytest.mark.asyncio +async def test_tool_client_includes_auth_header(): + """ToolClient sends Authorization header when api_key is provided.""" + captured_headers: dict[str, str] = {} + + def capture(request: Any, *args: Any, **kwargs: Any) -> Response: + captured_headers.update(dict(request.headers)) + return Response(200, json=SIMPLE_SPEC) + + with respx.mock: + respx.get("http://tools.local/openapi.json").mock(side_effect=capture) + client = ToolClient("http://tools.local", api_key="secret-token") + await client.get_spec() + + assert "authorization" in {k.lower() for k in captured_headers} + auth = next(v for k, v in captured_headers.items() if k.lower() == "authorization") + assert auth.startswith("Bearer ") + + +@pytest.mark.asyncio +async def test_tool_client_returns_plain_text_for_non_json(): + with respx.mock: + respx.get("http://tools.local/openapi.json").mock( + return_value=Response(200, json=SIMPLE_SPEC) + ) + respx.get("http://tools.local/items").mock( + return_value=Response(200, text="plain result", headers={"content-type": "text/plain"}) + ) + client = ToolClient("http://tools.local") + result = await client.call("list_items", {}) + assert result == "plain result" + + +# --------------------------------------------------------------------------- +# LLMClient.chat_with_tools +# --------------------------------------------------------------------------- + + +def _make_tool_response(tool_name: str, tool_call_id: str, args: str) -> MagicMock: + """Create a mock OpenAI response that requests a tool call.""" + tc = MagicMock() + tc.id = tool_call_id + tc.function.name = tool_name + tc.function.arguments = args + + msg = MagicMock() + msg.content = None + msg.tool_calls = [tc] + + choice = MagicMock() + choice.message = msg + + resp = MagicMock() + resp.choices = [choice] + return resp + + +def _make_text_response(content: str) -> MagicMock: + """Create a mock OpenAI response with plain text content.""" + msg = MagicMock() + msg.content = content + msg.tool_calls = None + + choice = MagicMock() + choice.message = msg + + resp = MagicMock() + resp.choices = [choice] + return resp + + +@pytest.mark.asyncio +async def test_chat_with_tools_no_tool_calls(): + """When the LLM returns no tool calls, reply is returned directly.""" + settings = _make_llm_settings() + llm = LLMClient(settings) + + mock_tool_client = AsyncMock() + mock_tool_client.get_tools = AsyncMock(return_value=[]) + + with patch.object( + llm._client.chat.completions, + "create", + new_callable=AsyncMock, + return_value=_make_text_response("Hello world"), + ): + result = await llm.chat_with_tools("Hi", mock_tool_client) + + assert result == "Hello world" + mock_tool_client.call.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_chat_with_tools_executes_tool_call(): + """Tool call is executed and result fed back; final text is returned.""" + settings = _make_llm_settings() + llm = LLMClient(settings) + + tool_response = _make_tool_response("list_items", "call-1", json.dumps({"limit": 5})) + final_response = _make_text_response("Here are the items: [foo, bar]") + + mock_tool_client = AsyncMock() + mock_tool_client.get_tools = AsyncMock(return_value=[{"type": "function"}]) + mock_tool_client.call = AsyncMock(return_value='[{"name": "foo"}, {"name": "bar"}]') + + with patch.object( + llm._client.chat.completions, + "create", + new_callable=AsyncMock, + side_effect=[tool_response, final_response], + ): + result = await llm.chat_with_tools("List items", mock_tool_client) + + assert result == "Here are the items: [foo, bar]" + mock_tool_client.call.assert_awaited_once_with("list_items", {"limit": 5}) + + +@pytest.mark.asyncio +async def test_chat_with_tools_handles_tool_error(): + """When tool execution fails the error is fed back to the LLM gracefully.""" + settings = _make_llm_settings() + llm = LLMClient(settings) + + tool_response = _make_tool_response("list_items", "call-1", json.dumps({})) + final_response = _make_text_response("I could not retrieve items due to an error.") + + mock_tool_client = AsyncMock() + mock_tool_client.get_tools = AsyncMock(return_value=[{"type": "function"}]) + mock_tool_client.call = AsyncMock(side_effect=RuntimeError("Server unreachable")) + + with patch.object( + llm._client.chat.completions, + "create", + new_callable=AsyncMock, + side_effect=[tool_response, final_response], + ) as mock_create: + result = await llm.chat_with_tools("List items", mock_tool_client) + + assert "error" in result.lower() or "could not" in result.lower() + # The second LLM call should include the tool error in messages + second_call_messages = mock_create.call_args_list[1].kwargs["messages"] + tool_result_msg = next(m for m in second_call_messages if m.get("role") == "tool") + assert "Tool call failed" in tool_result_msg["content"] + + +@pytest.mark.asyncio +async def test_chat_with_tools_respects_max_rounds(): + """chat_with_tools exits after max_tool_rounds even if LLM keeps calling tools.""" + settings = _make_llm_settings() + llm = LLMClient(settings) + + tool_response = _make_tool_response("list_items", "call-1", json.dumps({})) + + mock_tool_client = AsyncMock() + mock_tool_client.get_tools = AsyncMock(return_value=[{"type": "function"}]) + mock_tool_client.call = AsyncMock(return_value="[]") + + # Always returns a tool call – never a final answer + with patch.object( + llm._client.chat.completions, + "create", + new_callable=AsyncMock, + return_value=tool_response, + ) as mock_create: + await llm.chat_with_tools("List items", mock_tool_client, max_tool_rounds=3) + + assert mock_create.call_count == 3 + + +@pytest.mark.asyncio +async def test_chat_with_tools_passes_history(): + """History is included in the first LLM request.""" + settings = _make_llm_settings() + llm = LLMClient(settings) + + mock_tool_client = AsyncMock() + mock_tool_client.get_tools = AsyncMock(return_value=[]) + + history = [ + {"role": "user", "content": "previous"}, + {"role": "assistant", "content": "answer"}, + ] + + with patch.object( + llm._client.chat.completions, + "create", + new_callable=AsyncMock, + return_value=_make_text_response("reply"), + ) as mock_create: + await llm.chat_with_tools("new question", mock_tool_client, history=history) + + messages = mock_create.call_args.kwargs["messages"] + roles = [m["role"] for m in messages] + assert roles == ["system", "user", "assistant", "user"]