feat: MCP/OpenAPI tool calling support
This commit is contained in:
parent
a5037aa808
commit
76c6086db3
@ -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=
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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."""
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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"])
|
||||
|
||||
|
||||
1
steward/tools/__init__.py
Normal file
1
steward/tools/__init__.py
Normal file
@ -0,0 +1 @@
|
||||
"""MCP-compatible OpenAPI tool calling support."""
|
||||
259
steward/tools/client.py
Normal file
259
steward/tools/client.py
Normal 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
499
tests/test_tools.py
Normal 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"]
|
||||
Loading…
x
Reference in New Issue
Block a user