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