421a664336
- Remove legacy provider profiles, admin-token auth, /model command, model picker widget, and config.yaml LLM fields (design doc section 10) - Wire CLI/channels/cron and async sub-agents through the local snapshot entry; run creation rejects model config outside runtime_snapshot_id - Add periodic run-snapshot TTL cleanup to the config service lifespan - Isolate tests from the real config dir and activate the registry where run/model paths fail closed in bootstrap Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
368 lines
11 KiB
Python
368 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_auxiliary_chat_model")
|
|
@patch("EvoScientist.EvoScientist._ensure_chat_model")
|
|
@patch("EvoScientist.EvoScientist._ensure_config")
|
|
def test_default_middleware_includes_tool_selector(mock_config, mock_model, mock_aux, 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_auxiliary_chat_model")
|
|
@patch("EvoScientist.EvoScientist._ensure_chat_model")
|
|
@patch("EvoScientist.EvoScientist._ensure_config")
|
|
def test_tool_selector_ordering(mock_config, mock_model, mock_aux, 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
|