178 lines
5.4 KiB
Python
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)
|