"""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)