feat: MCP/OpenAPI tool calling support

This commit is contained in:
copilot-swe-agent[bot] 2026-07-25 12:50:39 +00:00 committed by GitHub
parent a5037aa808
commit 76c6086db3
8 changed files with 921 additions and 16 deletions

View File

@ -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=

View File

@ -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))

View File

@ -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."""

View File

@ -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,

View File

@ -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"])

View File

@ -0,0 +1 @@
"""MCP-compatible OpenAPI tool calling support."""

259
steward/tools/client.py Normal file
View File

@ -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")

499
tests/test_tools.py Normal file
View File

@ -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"]