feat: MCP/OpenAPI tool calling support
This commit is contained in:
committed by
GitHub
parent
a5037aa808
commit
76c6086db3
@@ -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"]
|
||||
Reference in New Issue
Block a user