Files
EvoScientist-Multi/tests/test_repetitive_tool_guard.py
T

178 lines
5.4 KiB
Python

"""Deterministic tool-loop guard and provider projection tests."""
from dataclasses import dataclass, replace
from typing import Any
import pytest
from langchain.agents.middleware.types import ModelResponse
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from EvoScientist.llm.errors import AgentControlError
from EvoScientist.middleware.repetitive_tool_guard import (
RepetitiveToolCallGuardMiddleware,
collapse_repetitive_tool_rounds,
)
def _round(
call_id: str,
*,
name: str = "execute",
command: str = "pwd",
content: str = "Error: invalid argument: command rejected by schema",
status: str = "error",
) -> list[Any]:
return [
AIMessage(
content="",
tool_calls=[{"id": call_id, "name": name, "args": {"command": command}}],
),
ToolMessage(
content=content,
tool_call_id=call_id,
name=name,
status=status,
),
]
@dataclass(frozen=True)
class _Request:
messages: list[Any]
tools: list[Any]
def override(self, **updates: Any):
return replace(self, **updates)
def test_provider_projection_keeps_first_and_last_deterministic_error_rounds():
messages = [HumanMessage(content="inspect")]
for index in range(4):
messages.extend(_round(f"call-{index}"))
messages.append(HumanMessage(content="continue"))
repair = collapse_repetitive_tool_rounds(messages, threshold=2)
assert repair.removed_rounds == 2
assert [m.type for m in repair.messages] == [
"human",
"ai",
"tool",
"ai",
"tool",
"human",
]
assert repair.messages[1].tool_calls[0]["id"] == "call-0"
assert repair.messages[3].tool_calls[0]["id"] == "call-3"
def test_successful_repeated_calls_are_never_projected_away():
messages = [
*_round("call-1", content="ok", status="success"),
*_round("call-2", content="ok", status="success"),
*_round("call-3", content="ok", status="success"),
]
repair = collapse_repetitive_tool_rounds(messages)
assert repair.messages == messages
assert repair.removed_rounds == 0
assert repair.tail_repetitions == 0
def test_transient_and_unknown_errors_do_not_count_as_semantic_loop():
transient = [
*_round("call-1", content="Error: connection timeout"),
*_round("call-2", content="Error: connection timeout"),
]
unknown = [
*_round("call-3", content="Error: something unusual"),
*_round("call-4", content="Error: something unusual"),
]
assert collapse_repetitive_tool_rounds(transient).tail_repetitions == 0
assert collapse_repetitive_tool_rounds(unknown).tail_consecutive_errors == 0
def test_generic_raw_execution_error_code_remains_unknown():
messages = _round("call-1", content="Error: something unusual")
messages[1].additional_kwargs["error_code"] = "TOOL_EXECUTION_FAILED"
repair = collapse_repetitive_tool_rounds(messages)
assert repair.tail_consecutive_errors == 0
def test_identical_tail_loop_stops_before_next_model_call():
request = _Request(
messages=[*_round("call-1"), *_round("call-2")],
tools=[{"name": "execute"}],
)
called = False
def handler(_request):
nonlocal called
called = True
return ModelResponse(result=[AIMessage(content="should not run")])
with pytest.raises(AgentControlError) as caught:
RepetitiveToolCallGuardMiddleware(threshold=2).wrap_model_call(request, handler)
assert caught.value.code == "MODEL_TOOL_LOOP_DETECTED"
assert called is False
def test_different_deterministic_errors_hit_consecutive_limit():
request = _Request(
messages=[
*_round("one", name="execute"),
*_round("two", name="read_file"),
*_round("three", name="search"),
],
tools=[],
)
with pytest.raises(AgentControlError) as caught:
RepetitiveToolCallGuardMiddleware(
threshold=0, max_consecutive_errors=3
).wrap_model_call(request, lambda _request: None)
assert caught.value.code == "MODEL_TOOL_ERROR_LIMIT"
def test_user_message_breaks_tail_loop_but_historical_projection_is_temporary():
original = [
*_round("call-1"),
*_round("call-2"),
*_round("call-3"),
HumanMessage(content="try a new approach"),
]
request = _Request(messages=original, tools=[])
captured = []
def handler(prepared):
captured.append(prepared)
return ModelResponse(result=[AIMessage(content="continued")])
RepetitiveToolCallGuardMiddleware().wrap_model_call(request, handler)
assert len(captured[0].messages) == 5
assert len(original) == 7
def test_zero_thresholds_disable_only_semantic_loop_guards():
request = _Request(messages=[*_round("one"), *_round("two")], tools=[])
captured = []
middleware = RepetitiveToolCallGuardMiddleware(
threshold=0, max_consecutive_errors=0
)
middleware.wrap_model_call(
request,
lambda prepared: (
captured.append(prepared) or ModelResponse(result=[AIMessage(content="ok")])
),
)
assert captured == [request]
@pytest.mark.parametrize("kwargs", [{"threshold": -1}, {"max_consecutive_errors": -1}])
def test_negative_threshold_is_rejected(kwargs):
with pytest.raises(ValueError, match="non-negative"):
RepetitiveToolCallGuardMiddleware(**kwargs)