563 lines
18 KiB
Python
563 lines
18 KiB
Python
"""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,
|
||
_discover_mcpo_server_names,
|
||
_operation_id,
|
||
_resolve_operation,
|
||
_schema_to_json_schema,
|
||
_spec_to_openai_tools,
|
||
)
|
||
from tests.conftest import make_settings
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Fixtures
|
||
# ---------------------------------------------------------------------------
|
||
|
||
MCPO_ROOT_SPEC: dict[str, Any] = {
|
||
"openapi": "3.1.0",
|
||
"info": {
|
||
"title": "MCP OpenAPI Proxy",
|
||
"version": "1.0",
|
||
"description": (
|
||
"Automatically generated API from MCP Tool Schemas\n\n"
|
||
"- **available tools**:\n"
|
||
" - [github](/github/docs)\n"
|
||
" - [memory](/memory/docs)\n"
|
||
),
|
||
},
|
||
"paths": {},
|
||
}
|
||
|
||
|
||
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:
|
||
"""Wrapper for test compatibility."""
|
||
defaults = dict(
|
||
telegram_bot_token="t",
|
||
openai_api_key="k",
|
||
openai_model="gpt-4o",
|
||
openai_system_prompt="You are Steward.",
|
||
)
|
||
defaults.update(kwargs)
|
||
return make_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
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def test_discover_mcpo_server_names():
|
||
assert _discover_mcpo_server_names(MCPO_ROOT_SPEC) == ["github", "memory"]
|
||
|
||
|
||
@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_get_tools_from_mcpo_subservers():
|
||
with respx.mock:
|
||
respx.get("http://tools.local/openapi.json").mock(
|
||
return_value=Response(200, json=MCPO_ROOT_SPEC)
|
||
)
|
||
respx.get("http://tools.local/github/openapi.json").mock(
|
||
return_value=Response(200, json=SIMPLE_SPEC)
|
||
)
|
||
respx.get("http://tools.local/memory/openapi.json").mock(
|
||
return_value=Response(200, json=SIMPLE_SPEC)
|
||
)
|
||
client = ToolClient("http://tools.local")
|
||
tools = await client.get_tools()
|
||
|
||
tool_names = {tool["function"]["name"] for tool in tools}
|
||
assert "github__list_items" in tool_names
|
||
assert "memory__list_items" in tool_names
|
||
assert len(tools) == 6
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_tool_client_call_mcpo_namespaced_tool():
|
||
with respx.mock:
|
||
respx.get("http://tools.local/openapi.json").mock(
|
||
return_value=Response(200, json=MCPO_ROOT_SPEC)
|
||
)
|
||
respx.get("http://tools.local/github/openapi.json").mock(
|
||
return_value=Response(200, json=SIMPLE_SPEC)
|
||
)
|
||
respx.get("http://tools.local/memory/openapi.json").mock(
|
||
return_value=Response(200, json=SIMPLE_SPEC)
|
||
)
|
||
respx.get("http://tools.local/github/items").mock(
|
||
return_value=Response(200, json=[{"id": 1}])
|
||
)
|
||
client = ToolClient("http://tools.local")
|
||
result = await client.call("github__list_items", {"limit": 10})
|
||
|
||
assert "id" in result
|
||
|
||
|
||
@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"]
|