Files
EvoScientist/tests/test_tool_selector_middleware.py
dinos b1dccf17ea fix(tool-selector): memory tools & state for main agent (#305)
* fix(tool-selector): always include memory tools

* fix(tool-selector): only track & show state for main agent

* docs: update docstring
2026-06-23 14:47:25 +01:00

366 lines
11 KiB
Python

"""Tests for LLMToolSelectorMiddleware integration."""
from typing import Any
from unittest.mock import MagicMock, patch
from langchain.agents.middleware.types import ModelRequest
from langchain_core.tools import BaseTool, StructuredTool
from EvoScientist.middleware.tool_selector import (
_ConditionalToolSelectorMiddleware,
_ToolSelectionTrackerMiddleware,
create_tool_selector_middleware,
)
def _tool(name: str) -> BaseTool:
def _func(value: str = "") -> str:
return value
return StructuredTool.from_function(
func=_func,
name=name,
description=f"{name} test tool",
)
def _request(tools: list[BaseTool | dict[str, Any]]) -> ModelRequest:
return ModelRequest(model=MagicMock(), messages=[], tools=tools)
def _mock_model():
"""Create a MagicMock model compatible with disable_thinking()."""
m = MagicMock(profile={"max_input_tokens": 200_000})
m.thinking = None
m.reasoning = None
return m
def _patched_create():
"""Create tool selector middleware without real LLM init."""
return [
_ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=MagicMock()),
threshold=20,
),
_ToolSelectionTrackerMiddleware(),
]
# Helper: patches needed to call create_tool_selector_middleware without LLM
def _factory_patches():
return (
patch(
"EvoScientist.middleware.tool_selector.disable_thinking",
return_value=MagicMock(),
create=True,
),
patch("EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()),
patch(
"langchain.agents.middleware.LLMToolSelectorMiddleware",
return_value=MagicMock(),
),
)
# ---------------------------------------------------------------------------
# Factory tests
# ---------------------------------------------------------------------------
def test_create_tool_selector_returns_list():
p1, p2, p3 = _factory_patches()
with p1, p2, p3:
result = create_tool_selector_middleware()
assert isinstance(result, list)
assert len(result) == 2
def test_create_tool_selector_always_include():
p1, p2, p3 = _factory_patches()
with p1, p2, p3 as mock_cls:
result = create_tool_selector_middleware(threshold=0)
request = _request(
[
_tool("think_tool"),
_tool("search_observations"),
_tool("read_memory"),
_tool("unrelated_tool"),
]
)
result[0].wrap_model_call(request, MagicMock())
mock_cls.assert_called_once()
call_kwargs = mock_cls.call_args[1]
assert call_kwargs["always_include"] == [
"read_memory",
"search_observations",
"think_tool",
]
def test_custom_threshold():
p1, p2, p3 = _factory_patches()
with p1, p2, p3:
result = create_tool_selector_middleware(threshold=5)
assert result[0]._threshold == 5
# ---------------------------------------------------------------------------
# Conditional + tracker unit tests
# ---------------------------------------------------------------------------
def test_conditional_skips_below_threshold():
"""When tools <= threshold, selector is skipped."""
mock_selector = MagicMock()
selector_factory = MagicMock(return_value=mock_selector)
cond = _ConditionalToolSelectorMiddleware(
selector_factory=selector_factory,
threshold=10,
)
request = MagicMock()
request.tools = [MagicMock() for _ in range(5)]
handler = MagicMock()
cond.wrap_model_call(request, handler)
handler.assert_called_once_with(request)
selector_factory.assert_not_called()
mock_selector.wrap_model_call.assert_not_called()
def test_conditional_runs_above_threshold():
"""When tools > threshold, selector runs."""
mock_selector = MagicMock()
selector_factory = MagicMock(return_value=mock_selector)
cond = _ConditionalToolSelectorMiddleware(
selector_factory=selector_factory,
threshold=10,
)
request = MagicMock()
request.tools = [MagicMock() for _ in range(15)]
handler = MagicMock()
cond.wrap_model_call(request, handler)
selector_factory.assert_called_once_with([])
mock_selector.wrap_model_call.assert_called_once()
handler.assert_not_called()
def test_selector_active_flag():
"""_selector_active flag is True during selection, False after."""
import EvoScientist.middleware.tool_selector as ts_mod
mock_selector = MagicMock()
def fake_selector_call(request, handler):
assert ts_mod._selector_active is True
return handler(request)
mock_selector.wrap_model_call.side_effect = fake_selector_call
cond = _ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=mock_selector),
threshold=5,
)
request = MagicMock()
request.tools = [MagicMock() for _ in range(10)]
handler = MagicMock()
cond.wrap_model_call(request, handler)
assert ts_mod._selector_active is False
def test_selector_can_disable_stream_tracking():
"""Selection can run without touching the main-agent stream/UI globals."""
import EvoScientist.middleware.tool_selector as ts_mod
mock_selector = MagicMock()
def fake_selector_call(request, handler):
assert ts_mod._selector_active is False
return handler(request)
mock_selector.wrap_model_call.side_effect = fake_selector_call
cond = _ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=mock_selector),
threshold=5,
track_stream_selection=False,
)
ts_mod._total_tools_count = 99
request = MagicMock()
request.tools = [MagicMock() for _ in range(10)]
handler = MagicMock()
cond.wrap_model_call(request, handler)
mock_selector.wrap_model_call.assert_called_once()
handler.assert_called_once()
assert ts_mod._selector_active is False
assert ts_mod._total_tools_count == 99
def test_selector_always_includes_available_memory_tools():
"""Adaptive selection must mark available memory tools as mandatory."""
calls = []
class FakeSelector:
def __init__(self, always_include):
self.always_include = always_include
def wrap_model_call(self, request, handler):
calls.append(self.always_include)
return handler(request)
def selector_factory(always_include):
return FakeSelector(always_include)
request = _request(
[
_tool("think_tool"),
_tool("search_observations"),
_tool("read_memory"),
_tool("unrelated_tool"),
]
)
cond = _ConditionalToolSelectorMiddleware(
selector_factory=selector_factory,
threshold=0,
always_include=frozenset(
{
"think_tool",
"task",
"search_observations",
"read_memory",
"record_observation",
}
),
)
handler = MagicMock()
cond.wrap_model_call(request, handler)
assert calls == [
[
"read_memory",
"search_observations",
"think_tool",
]
]
def test_selector_resolved_once_across_repeated_requests():
"""Agent tools are stable, so build the selector once and reuse it."""
mock_selector = MagicMock()
mock_selector.wrap_model_call.side_effect = lambda request, handler: handler(
request
)
selector_factory = MagicMock(return_value=mock_selector)
cond = _ConditionalToolSelectorMiddleware(
selector_factory=selector_factory,
threshold=0,
always_include=frozenset({"think_tool", "search_observations"}),
)
tools = [
_tool("think_tool"),
_tool("search_observations"),
_tool("unrelated_tool"),
]
for _ in range(3):
cond.wrap_model_call(_request(tools), MagicMock())
selector_factory.assert_called_once_with(["search_observations", "think_tool"])
assert mock_selector.wrap_model_call.call_count == 3
def test_tracker_captures_tools():
"""Tracker middleware captures tool names from request."""
tracker = _ToolSelectionTrackerMiddleware()
tool1 = _tool("read_file")
tool2 = _tool("execute")
request = MagicMock()
request.tools = [tool1, tool2]
handler = MagicMock()
tracker.wrap_model_call(request, handler)
handler.assert_called_once_with(request)
import EvoScientist.middleware.tool_selector as ts_mod
assert ts_mod._current_selected_tools == ["read_file", "execute"]
# ---------------------------------------------------------------------------
# Integration tests
# ---------------------------------------------------------------------------
@patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=lambda *a, **kw: _patched_create(),
)
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
def test_default_middleware_includes_tool_selector(mock_config, mock_model, mock_ts):
mock_model.return_value = _mock_model()
cfg = MagicMock()
cfg.enable_ask_user = False
cfg.auto_approve = False
cfg.auxiliary_model = ""
cfg.auxiliary_provider = ""
mock_config.return_value = cfg
from EvoScientist.EvoScientist import _get_default_middleware
mw = _get_default_middleware()
type_names = [type(m).__name__ for m in mw]
assert "_ConditionalToolSelectorMiddleware" in type_names
assert "_ToolSelectionTrackerMiddleware" in type_names
@patch("EvoScientist.EvoScientist._ensure_chat_model")
def test_subagent_no_tool_selector(mock_model):
mock_model.return_value = _mock_model()
from EvoScientist.EvoScientist import _inject_subagent_middleware
subs = [{"name": "test-agent"}]
_inject_subagent_middleware(subs)
type_names = [type(m).__name__ for m in subs[0]["middleware"]]
assert "_ConditionalToolSelectorMiddleware" not in type_names
@patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=lambda *a, **kw: _patched_create(),
)
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
def test_tool_selector_ordering(mock_config, mock_model, mock_ts):
"""ToolSelector should come after ToolErrorHandler and before Memory."""
mock_model.return_value = _mock_model()
cfg = MagicMock()
cfg.enable_ask_user = False
cfg.auto_approve = False
cfg.auxiliary_model = ""
cfg.auxiliary_provider = ""
mock_config.return_value = cfg
from EvoScientist.EvoScientist import _get_default_middleware
mw = _get_default_middleware()
type_names = [type(m).__name__ for m in mw]
ts_idx = type_names.index("_ConditionalToolSelectorMiddleware")
tracker_idx = type_names.index("_ToolSelectionTrackerMiddleware")
te_idx = type_names.index("ToolErrorHandlerMiddleware")
mem_idx = type_names.index("EvoMemoryMiddleware")
assert te_idx < ts_idx < tracker_idx < mem_idx