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
This commit is contained in:
@@ -724,7 +724,10 @@ def _get_default_middleware(
|
||||
ModelFallbackMiddleware(),
|
||||
ContextOverflowMapperMiddleware(),
|
||||
ToolErrorHandlerMiddleware(),
|
||||
*create_tool_selector_middleware(model=tool_selector_model),
|
||||
*create_tool_selector_middleware(
|
||||
model=tool_selector_model,
|
||||
track_stream_selection=not for_async_subagent,
|
||||
),
|
||||
# Interpreter prompt must land before runtime/memory context, so this
|
||||
# middleware sits ahead of runtime_context in the stack.
|
||||
create_code_interpreter_middleware(
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""LLMToolSelectorMiddleware configuration for EvoScientist.
|
||||
|
||||
Wraps LangChain's built-in ``LLMToolSelectorMiddleware`` with project-specific
|
||||
defaults and a tracker that captures which tools were selected.
|
||||
defaults and an optional stream tracker that captures which tools were selected.
|
||||
|
||||
The selector only activates when the agent has more than ``threshold`` tools
|
||||
(default 20). Below that, the extra LLM call isn't worth the token savings.
|
||||
@@ -16,19 +16,23 @@ Usage::
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
AIMessage,
|
||||
ExtendedModelResponse,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
)
|
||||
from langchain_core.language_models import BaseChatModel
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Module-level storage for tool selection state.
|
||||
# Updated by _ToolSelectionTrackerMiddleware; read by stream/events.py.
|
||||
# Module-level storage for main-agent tool-selection UI state.
|
||||
# Updated only when stream tracking is enabled; read by stream/events.py.
|
||||
_current_selected_tools: list[str] = []
|
||||
_last_emitted_tools: list[str] = [] # last selection shown to user
|
||||
_total_tools_count: int = 0 # total tools before selection
|
||||
@@ -37,6 +41,31 @@ _selector_active: bool = False
|
||||
# Default threshold: only run tool selection when tools exceed this count.
|
||||
# Base tools are ~14; selector activates when MCP tools push count above 26.
|
||||
DEFAULT_TOOL_THRESHOLD = 26
|
||||
DEFAULT_ALWAYS_INCLUDE_TOOLS: frozenset[str] = frozenset(
|
||||
{
|
||||
"think_tool",
|
||||
"task",
|
||||
"read_memory",
|
||||
"record_observation",
|
||||
"search_observations",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _tool_name(tool: BaseTool | dict[str, Any]) -> str | None:
|
||||
if isinstance(tool, BaseTool):
|
||||
return tool.name or None
|
||||
name = tool.get("name")
|
||||
return name if isinstance(name, str) and name else None
|
||||
|
||||
|
||||
def _available_always_include(
|
||||
tools: Iterable[BaseTool | dict[str, Any]],
|
||||
candidates: frozenset[str],
|
||||
) -> list[str]:
|
||||
"""Return mandatory BaseTool names that exist on this request."""
|
||||
available_names = {tn for tool in tools if (tn := _tool_name(tool))}
|
||||
return sorted(candidates & available_names)
|
||||
|
||||
|
||||
class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
||||
@@ -45,30 +74,47 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
||||
Skips the selection LLM call when ``len(request.tools) <= threshold``,
|
||||
avoiding unnecessary overhead for agents with few tools.
|
||||
|
||||
Sets ``_selector_active`` flag during the selector's internal LLM call
|
||||
so the streaming layer can suppress its output.
|
||||
When stream tracking is enabled, sets ``_selector_active`` during the
|
||||
selector's internal LLM call so the streaming layer can suppress its output.
|
||||
"""
|
||||
|
||||
name = "conditional_tool_selector"
|
||||
|
||||
def __init__(
|
||||
self, selector: AgentMiddleware, threshold: int = DEFAULT_TOOL_THRESHOLD
|
||||
self,
|
||||
selector_factory: Callable[[list[str]], AgentMiddleware],
|
||||
threshold: int = DEFAULT_TOOL_THRESHOLD,
|
||||
*,
|
||||
always_include: frozenset[str] | None = None,
|
||||
track_stream_selection: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self._selector = selector
|
||||
self._selector_factory = selector_factory
|
||||
self._threshold = threshold
|
||||
self._always_include = always_include or frozenset()
|
||||
self._track_stream_selection = track_stream_selection
|
||||
# Agent tools are fixed after graph construction, so the filtered
|
||||
# always-include set is stable for this middleware instance.
|
||||
self._selector: AgentMiddleware | None = None
|
||||
|
||||
def _build_selector(self, request: ModelRequest) -> AgentMiddleware:
|
||||
if self._selector is None:
|
||||
names = _available_always_include(request.tools, self._always_include)
|
||||
self._selector = self._selector_factory(names)
|
||||
return self._selector
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
if len(request.tools or []) <= self._threshold:
|
||||
) -> ModelResponse | AIMessage | ExtendedModelResponse:
|
||||
if len(request.tools) <= self._threshold:
|
||||
return handler(request)
|
||||
|
||||
global _selector_active, _total_tools_count
|
||||
_selector_active = True
|
||||
_total_tools_count = len(request.tools or [])
|
||||
if self._track_stream_selection:
|
||||
global _selector_active, _total_tools_count
|
||||
_selector_active = True
|
||||
_total_tools_count = len(request.tools)
|
||||
|
||||
# Track whether handler was called — if so, any exception is from
|
||||
# the downstream model, not the selector, and must propagate.
|
||||
@@ -76,56 +122,65 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
||||
|
||||
def _handler_after_selection(req: ModelRequest) -> ModelResponse:
|
||||
nonlocal _handler_called
|
||||
global _selector_active
|
||||
_handler_called = True
|
||||
_selector_active = False
|
||||
if self._track_stream_selection:
|
||||
global _selector_active
|
||||
_selector_active = False
|
||||
return handler(req)
|
||||
|
||||
try:
|
||||
return self._selector.wrap_model_call(request, _handler_after_selection)
|
||||
return self._build_selector(request).wrap_model_call(
|
||||
request, _handler_after_selection
|
||||
)
|
||||
except Exception:
|
||||
if _handler_called:
|
||||
raise # Error from downstream model — don't retry
|
||||
# Selector itself failed (e.g., structured output not supported).
|
||||
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
||||
_selector_active = False
|
||||
if self._track_stream_selection:
|
||||
_selector_active = False
|
||||
return handler(request)
|
||||
finally:
|
||||
_selector_active = False
|
||||
if self._track_stream_selection:
|
||||
_selector_active = False
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
if len(request.tools or []) <= self._threshold:
|
||||
) -> ModelResponse | AIMessage | ExtendedModelResponse:
|
||||
if len(request.tools) <= self._threshold:
|
||||
return await handler(request)
|
||||
|
||||
global _selector_active, _total_tools_count
|
||||
_selector_active = True
|
||||
_total_tools_count = len(request.tools or [])
|
||||
if self._track_stream_selection:
|
||||
global _selector_active, _total_tools_count
|
||||
_selector_active = True
|
||||
_total_tools_count = len(request.tools)
|
||||
|
||||
_handler_called = False
|
||||
|
||||
async def _handler_after_selection(req: ModelRequest) -> ModelResponse:
|
||||
nonlocal _handler_called
|
||||
global _selector_active
|
||||
_handler_called = True
|
||||
_selector_active = False
|
||||
if self._track_stream_selection:
|
||||
global _selector_active
|
||||
_selector_active = False
|
||||
return await handler(req)
|
||||
|
||||
try:
|
||||
return await self._selector.awrap_model_call(
|
||||
return await self._build_selector(request).awrap_model_call(
|
||||
request, _handler_after_selection
|
||||
)
|
||||
except Exception:
|
||||
if _handler_called:
|
||||
raise
|
||||
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
||||
_selector_active = False
|
||||
if self._track_stream_selection:
|
||||
_selector_active = False
|
||||
return await handler(request)
|
||||
finally:
|
||||
_selector_active = False
|
||||
if self._track_stream_selection:
|
||||
_selector_active = False
|
||||
|
||||
|
||||
class _ToolSelectionTrackerMiddleware(AgentMiddleware):
|
||||
@@ -144,7 +199,7 @@ class _ToolSelectionTrackerMiddleware(AgentMiddleware):
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
global _current_selected_tools
|
||||
tools = [t.name for t in request.tools if hasattr(t, "name")]
|
||||
tools = [name for tool in request.tools if (name := _tool_name(tool))]
|
||||
_current_selected_tools = tools
|
||||
if tools:
|
||||
logger.debug("Selected tools: %s", tools)
|
||||
@@ -156,7 +211,7 @@ class _ToolSelectionTrackerMiddleware(AgentMiddleware):
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
global _current_selected_tools
|
||||
tools = [t.name for t in request.tools if hasattr(t, "name")]
|
||||
tools = [name for tool in request.tools if (name := _tool_name(tool))]
|
||||
_current_selected_tools = tools
|
||||
if tools:
|
||||
logger.debug("Selected tools: %s", tools)
|
||||
@@ -167,25 +222,32 @@ def create_tool_selector_middleware(
|
||||
threshold: int = DEFAULT_TOOL_THRESHOLD,
|
||||
*,
|
||||
model: BaseChatModel | None = None,
|
||||
track_stream_selection: bool = True,
|
||||
):
|
||||
"""Build LLMToolSelectorMiddleware + tracker with EvoScientist defaults.
|
||||
|
||||
Returns a list of two middleware:
|
||||
Returns middleware for adaptive tool selection:
|
||||
1. Conditional wrapper around ``LLMToolSelectorMiddleware`` — only
|
||||
activates when ``len(tools) > threshold``
|
||||
2. ``_ToolSelectionTrackerMiddleware`` — captures selected tool names
|
||||
2. Optional ``_ToolSelectionTrackerMiddleware`` — captures selected tool
|
||||
names for the main-agent stream UI when ``track_stream_selection`` is true
|
||||
|
||||
Args:
|
||||
model: Chat model for tool selection. If *None*, the default
|
||||
model is resolved via ``_ensure_chat_model()``.
|
||||
threshold: Minimum number of tools to trigger selection.
|
||||
Default 20. Set to 0 to always run selection.
|
||||
Default 26. Set to 0 to always run selection.
|
||||
track_stream_selection: Whether to update process-global stream/UI
|
||||
state. Disable for async sub-agents that should still select tools
|
||||
but should not drive the main-agent tool-selection widget.
|
||||
|
||||
``think_tool`` and ``task`` are always included because:
|
||||
``think_tool``, ``task``, and memory tools are always included because:
|
||||
|
||||
- ``think_tool``: required every step for structured reflection
|
||||
- ``task``: core delegation mechanism; tested and confirmed the selector
|
||||
model never auto-selects it (0/5 complex queries)
|
||||
- memory tools: referenced by memory prompts; filtering them makes the
|
||||
agent unable to use memory even when the prompt tells it to
|
||||
"""
|
||||
from langchain.agents.middleware import LLMToolSelectorMiddleware
|
||||
|
||||
@@ -197,20 +259,47 @@ def create_tool_selector_middleware(
|
||||
model = _ensure_chat_model()
|
||||
safe_model = disable_thinking(model)
|
||||
|
||||
selector = LLMToolSelectorMiddleware(
|
||||
model=safe_model,
|
||||
system_prompt=(
|
||||
"You are selecting tools for a scientific research agent. "
|
||||
"Tasks often involve multi-step workflows. "
|
||||
"Select tools that cover both the immediate need and "
|
||||
"likely follow-up steps. "
|
||||
"If the query is broad or all tools seem relevant, "
|
||||
"select all of them — filtering is not always necessary."
|
||||
),
|
||||
always_include=["think_tool", "task"],
|
||||
system_prompt = (
|
||||
"You are selecting tools for a scientific research agent. "
|
||||
"Tasks often involve multi-step workflows. "
|
||||
"Select tools that cover both the immediate need and "
|
||||
"likely follow-up steps. "
|
||||
"If the query is broad or all tools seem relevant, "
|
||||
"select all of them — filtering is not always necessary."
|
||||
)
|
||||
|
||||
return [
|
||||
_ConditionalToolSelectorMiddleware(selector, threshold=threshold),
|
||||
_ToolSelectionTrackerMiddleware(),
|
||||
def selector_factory(always_include: list[str]) -> AgentMiddleware:
|
||||
return LLMToolSelectorMiddleware(
|
||||
model=safe_model,
|
||||
system_prompt=system_prompt,
|
||||
always_include=always_include,
|
||||
)
|
||||
|
||||
middleware: list[AgentMiddleware] = [
|
||||
_ConditionalToolSelectorMiddleware(
|
||||
selector_factory=selector_factory,
|
||||
threshold=threshold,
|
||||
always_include=DEFAULT_ALWAYS_INCLUDE_TOOLS,
|
||||
track_stream_selection=track_stream_selection,
|
||||
),
|
||||
]
|
||||
if track_stream_selection:
|
||||
middleware.append(_ToolSelectionTrackerMiddleware())
|
||||
return middleware
|
||||
|
||||
|
||||
def reset_tool_selection_state_for_tests() -> None:
|
||||
"""Reset the process-global tool-selection state.
|
||||
|
||||
The selector/tracker record the last selected tools and the selector-active
|
||||
flag in module globals that ``stream/tool_selection.py`` reads to suppress
|
||||
selector chatter. Tests that drive the selector must not leak that state
|
||||
into later tests; an autouse fixture resets it around every test.
|
||||
"""
|
||||
global _current_selected_tools, _last_emitted_tools
|
||||
global _total_tools_count, _selector_active
|
||||
|
||||
_current_selected_tools = []
|
||||
_last_emitted_tools = []
|
||||
_total_tools_count = 0
|
||||
_selector_active = False
|
||||
|
||||
@@ -710,7 +710,11 @@ class _V3EventProcessor:
|
||||
) -> list[dict[str, Any]]:
|
||||
if not text or subagent is not None:
|
||||
return []
|
||||
return [self.emitter.thinking(text).data]
|
||||
suppressed, events = self._selector.process_thinking(text)
|
||||
if suppressed:
|
||||
return events
|
||||
events.append(self.emitter.thinking(text).data)
|
||||
return events
|
||||
|
||||
def _emit_summarization_text(self, text: str) -> list[dict[str, Any]]:
|
||||
if not text:
|
||||
|
||||
@@ -63,6 +63,16 @@ class _ToolSelectionSuppressor:
|
||||
|
||||
return False, events, text
|
||||
|
||||
def process_thinking(self, text: str) -> tuple[bool, list[dict[str, Any]]]:
|
||||
"""Suppress selector-model reasoning while preserving pending UI events."""
|
||||
events = self._emit_selection_if_ready(text)
|
||||
if not text:
|
||||
return False, events
|
||||
if self._selector_call_active():
|
||||
self._was_active = True
|
||||
return True, events
|
||||
return False, events
|
||||
|
||||
@staticmethod
|
||||
def _json_buffer_kind(text: str) -> str:
|
||||
try:
|
||||
@@ -81,11 +91,19 @@ class _ToolSelectionSuppressor:
|
||||
def _selector_context_active(self) -> bool:
|
||||
if self._was_active:
|
||||
return True
|
||||
return self._selector_call_active() or self._selection_pending()
|
||||
|
||||
@staticmethod
|
||||
def _selector_call_active() -> bool:
|
||||
import EvoScientist.middleware.tool_selector as selector_mod
|
||||
|
||||
return bool(
|
||||
selector_mod._selector_active or selector_mod._current_selected_tools
|
||||
)
|
||||
return bool(selector_mod._selector_active)
|
||||
|
||||
@staticmethod
|
||||
def _selection_pending() -> bool:
|
||||
import EvoScientist.middleware.tool_selector as selector_mod
|
||||
|
||||
return bool(selector_mod._current_selected_tools)
|
||||
|
||||
def flush_selection(self) -> list[dict[str, Any]]:
|
||||
return self._emit_selection_if_ready("")
|
||||
|
||||
@@ -31,6 +31,26 @@ def run_async_fixture():
|
||||
return run_async
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_tool_selection_state():
|
||||
"""Isolate the process-global tool-selection state around every test.
|
||||
|
||||
``middleware.tool_selector`` records the last selected tools and the
|
||||
selector-active flag in module globals that ``stream/tool_selection.py``
|
||||
reads to decide whether to suppress selector output. A test that drives the
|
||||
selector or tracker would otherwise leave those globals set and silently
|
||||
flip unrelated streaming tests later in the same process. Reset on both ends
|
||||
so order and worker sharding can't reintroduce the leak.
|
||||
"""
|
||||
from EvoScientist.middleware.tool_selector import (
|
||||
reset_tool_selection_state_for_tests,
|
||||
)
|
||||
|
||||
reset_tool_selection_state_for_tests()
|
||||
yield
|
||||
reset_tool_selection_state_for_tests()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_tool_call():
|
||||
"""A minimal tool call dict."""
|
||||
|
||||
@@ -283,3 +283,35 @@ def test_async_subagent_mode_filters_ask_user(
|
||||
"AskUserMiddleware leaked into async sub-agent middleware — its "
|
||||
"interrupt() call deadlocks the deployed graph (no UI to resume)."
|
||||
)
|
||||
|
||||
|
||||
@patch(
|
||||
"EvoScientist.middleware.create_tool_selector_middleware",
|
||||
return_value=[MagicMock()],
|
||||
)
|
||||
@patch("EvoScientist.EvoScientist._ensure_chat_model")
|
||||
@patch("EvoScientist.EvoScientist._ensure_config")
|
||||
def test_async_subagent_disables_tool_selector_stream_tracking(
|
||||
mock_config, mock_chat, mock_tool_selector
|
||||
):
|
||||
"""Async subagents still select tools, but must not drive main-agent UI state."""
|
||||
cfg = MagicMock()
|
||||
cfg.enable_ask_user = False
|
||||
cfg.auto_mode = False
|
||||
cfg.auto_approve = False
|
||||
cfg.model_fallbacks = None
|
||||
cfg.memory_profile_enabled = True
|
||||
cfg.memory_observations_enabled = True
|
||||
cfg.memory_observation_writer = MemoryObservationWriter.ALL
|
||||
cfg.memory_workers_enabled = True
|
||||
cfg.auxiliary_model = ""
|
||||
cfg.auxiliary_provider = ""
|
||||
mock_config.return_value = cfg
|
||||
mock_chat.return_value = MagicMock(profile={"max_input_tokens": 200_000})
|
||||
|
||||
from EvoScientist.EvoScientist import _get_default_middleware
|
||||
|
||||
_get_default_middleware(for_async_subagent=True)
|
||||
|
||||
mock_tool_selector.assert_called_once()
|
||||
assert mock_tool_selector.call_args.kwargs["track_stream_selection"] is False
|
||||
|
||||
@@ -407,6 +407,61 @@ class TestV3ProtocolStreaming:
|
||||
assert len(thinking_events) == 1
|
||||
assert thinking_events[0]["content"] == "Think once."
|
||||
|
||||
def test_tool_selector_reasoning_delta_is_suppressed(self):
|
||||
"""Selector reasoning must not appear as main-agent thinking."""
|
||||
import EvoScientist.middleware.tool_selector as selector_mod
|
||||
|
||||
original_active = selector_mod._selector_active
|
||||
selector_mod._selector_active = True
|
||||
try:
|
||||
agent = FakeV3Agent(
|
||||
[
|
||||
protocol_event(
|
||||
"messages",
|
||||
(
|
||||
{
|
||||
"event": "content-block-delta",
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"type": "reasoning-delta",
|
||||
"reasoning": "selector-only thought",
|
||||
},
|
||||
},
|
||||
{},
|
||||
),
|
||||
)
|
||||
]
|
||||
)
|
||||
events = collect_events(agent)
|
||||
finally:
|
||||
selector_mod._selector_active = original_active
|
||||
|
||||
assert not any(
|
||||
e.get("type") == "thinking" and e.get("content") == "selector-only thought"
|
||||
for e in events
|
||||
)
|
||||
|
||||
def test_tool_selector_whole_message_reasoning_is_suppressed(self):
|
||||
"""Selector reasoning in whole-message payloads is also hidden."""
|
||||
import EvoScientist.middleware.tool_selector as selector_mod
|
||||
|
||||
original_active = selector_mod._selector_active
|
||||
selector_mod._selector_active = True
|
||||
try:
|
||||
message = AIMessage(
|
||||
additional_kwargs={"reasoning_content": "selector whole thought"},
|
||||
content="",
|
||||
)
|
||||
agent = FakeV3Agent([protocol_event("messages", (message, {}))])
|
||||
events = collect_events(agent)
|
||||
finally:
|
||||
selector_mod._selector_active = original_active
|
||||
|
||||
assert not any(
|
||||
e.get("type") == "thinking" and e.get("content") == "selector whole thought"
|
||||
for e in events
|
||||
)
|
||||
|
||||
def test_tool_events_emit_call_and_result(self):
|
||||
"""v3 tool projection events become UI tool call/result events."""
|
||||
output = ToolMessage(
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
"""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,
|
||||
@@ -9,6 +13,21 @@ from EvoScientist.middleware.tool_selector import (
|
||||
)
|
||||
|
||||
|
||||
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})
|
||||
@@ -20,7 +39,10 @@ def _mock_model():
|
||||
def _patched_create():
|
||||
"""Create tool selector middleware without real LLM init."""
|
||||
return [
|
||||
_ConditionalToolSelectorMiddleware(MagicMock(), threshold=20),
|
||||
_ConditionalToolSelectorMiddleware(
|
||||
selector_factory=MagicMock(return_value=MagicMock()),
|
||||
threshold=20,
|
||||
),
|
||||
_ToolSelectionTrackerMiddleware(),
|
||||
]
|
||||
|
||||
@@ -57,11 +79,23 @@ def test_create_tool_selector_returns_list():
|
||||
def test_create_tool_selector_always_include():
|
||||
p1, p2, p3 = _factory_patches()
|
||||
with p1, p2, p3 as mock_cls:
|
||||
create_tool_selector_middleware()
|
||||
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 "think_tool" in call_kwargs["always_include"]
|
||||
assert "task" in call_kwargs["always_include"]
|
||||
assert call_kwargs["always_include"] == [
|
||||
"read_memory",
|
||||
"search_observations",
|
||||
"think_tool",
|
||||
]
|
||||
|
||||
|
||||
def test_custom_threshold():
|
||||
@@ -79,7 +113,11 @@ def test_custom_threshold():
|
||||
def test_conditional_skips_below_threshold():
|
||||
"""When tools <= threshold, selector is skipped."""
|
||||
mock_selector = MagicMock()
|
||||
cond = _ConditionalToolSelectorMiddleware(mock_selector, threshold=10)
|
||||
selector_factory = MagicMock(return_value=mock_selector)
|
||||
cond = _ConditionalToolSelectorMiddleware(
|
||||
selector_factory=selector_factory,
|
||||
threshold=10,
|
||||
)
|
||||
|
||||
request = MagicMock()
|
||||
request.tools = [MagicMock() for _ in range(5)]
|
||||
@@ -87,19 +125,25 @@ def test_conditional_skips_below_threshold():
|
||||
|
||||
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()
|
||||
cond = _ConditionalToolSelectorMiddleware(mock_selector, threshold=10)
|
||||
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()
|
||||
|
||||
@@ -115,7 +159,10 @@ def test_selector_active_flag():
|
||||
return handler(request)
|
||||
|
||||
mock_selector.wrap_model_call.side_effect = fake_selector_call
|
||||
cond = _ConditionalToolSelectorMiddleware(mock_selector, threshold=5)
|
||||
cond = _ConditionalToolSelectorMiddleware(
|
||||
selector_factory=MagicMock(return_value=mock_selector),
|
||||
threshold=5,
|
||||
)
|
||||
|
||||
request = MagicMock()
|
||||
request.tools = [MagicMock() for _ in range(10)]
|
||||
@@ -125,16 +172,120 @@ def test_selector_active_flag():
|
||||
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()
|
||||
mock_tool1 = MagicMock()
|
||||
mock_tool1.name = "read_file"
|
||||
mock_tool2 = MagicMock()
|
||||
mock_tool2.name = "execute"
|
||||
tool1 = _tool("read_file")
|
||||
tool2 = _tool("execute")
|
||||
|
||||
request = MagicMock()
|
||||
request.tools = [mock_tool1, mock_tool2]
|
||||
request.tools = [tool1, tool2]
|
||||
handler = MagicMock()
|
||||
|
||||
tracker.wrap_model_call(request, handler)
|
||||
|
||||
Reference in New Issue
Block a user