from __future__ import annotations import pytest from langchain.agents.middleware.types import ModelRequest from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage from EvoScientist.llm.contracts import EvoRuntimeError from EvoScientist.llm.errors import ( AgentControlError, ModelProviderResponseError, ModelToolProtocolError, ) from EvoScientist.middleware.evo_route_fallback import EvoRouteFallbackMiddleware class _Model: def __init__(self, route_key: str, *, supports_tools: bool | None = None) -> None: self.metadata = {"route_key": route_key} if supports_tools is not None: self.metadata["route_supports_tools"] = supports_tools class _Health: def __init__(self, open_routes: set[str] | None = None) -> None: self.open_routes = open_routes or set() def is_open(self, route_key: str) -> bool: return route_key in self.open_routes def _request(model: _Model) -> ModelRequest: return ModelRequest(model=model, messages=[], tools=[]) def test_route_strips_tools_when_frozen_model_capability_is_false(): model = _Model("text-only", supports_tools=False) request = ModelRequest( model=model, messages=[], tools=[{"type": "function", "function": {"name": "search"}}], ) def handler(routed: ModelRequest): return routed.tools assert EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) == [] def test_tools_disabled_route_removes_checkpoint_tool_protocol(): model = _Model("text-only", supports_tools=False) request = ModelRequest( model=model, messages=[ SystemMessage("Keep the answer concise."), AIMessage( content="", tool_calls=[{"name": "search", "args": {"q": "K3"}, "id": "call_1"}], ), ToolMessage(content="tool result", tool_call_id="call_1", name="search"), AIMessage( content="The tool found a result.", additional_kwargs={"tool_calls": [{"id": "call_2"}]}, ), HumanMessage("Continue."), ], tools=[{"type": "function", "function": {"name": "search"}}], tool_choice="required", response_format={"type": "json_object"}, model_settings={"temperature": 0, "max_tokens": 1}, ) def handler(routed: ModelRequest): return routed routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) assert routed.tools == [] assert routed.tool_choice is None assert routed.response_format is None assert routed.model_settings == {} assert [message.type for message in routed.messages] == [ "system", "ai", "ai", "human", ] assistant = routed.messages[1] assert isinstance(assistant, AIMessage) assert assistant.content == "[Completed tool result: search]\ntool result" assert routed.messages[2].content == "The tool found a result." assert routed.messages[2].tool_calls == [] assert "tool_calls" not in routed.messages[2].additional_kwargs def test_tools_disabled_route_projects_tool_result_and_content_blocks_to_text(): model = _Model("text-only", supports_tools=False) request = ModelRequest( model=model, messages=[ AIMessage( content=[ {"type": "text", "text": "I checked the source."}, {"type": "tool_call", "id": "call_1", "name": "read"}, ], ), ToolMessage( content="x" * 12_100, tool_call_id="call_1", name="read", ), ], tools=[{"type": "function", "function": {"name": "read"}}], ) routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, lambda item: item) assert routed.tools == [] assert len(routed.messages) == 2 assert routed.messages[0].content == "I checked the source." assert routed.messages[0].tool_calls == [] assert routed.messages[1].content.startswith("[Completed tool result: read]\n") assert routed.messages[1].content.endswith("[Tool result truncated]") def test_route_keeps_tools_when_frozen_model_capability_is_true(): model = _Model("tools", supports_tools=True) tools = [{"type": "function", "function": {"name": "search"}}] request = ModelRequest( model=model, messages=[], tools=tools, model_settings={"temperature": 0}, ) def handler(routed: ModelRequest): return routed routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) assert routed.tools == tools assert routed.model_settings == {} def test_route_drops_empty_assistant_history_but_keeps_tool_calls(): model = _Model("tools", supports_tools=True) tool_message = AIMessage( content="", tool_calls=[{"name": "search", "args": {"q": "K3"}, "id": "call_1"}], ) request = ModelRequest( model=model, messages=[ HumanMessage("First"), AIMessage(content="", additional_kwargs={"reasoning_content": "hidden"}), HumanMessage("Continue"), tool_message, ], tools=[], ) routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, lambda item: item) assert routed.messages == [request.messages[0], request.messages[2], tool_message] @pytest.mark.asyncio async def test_empty_provider_response_retries_same_route_once(): primary = _Model("primary", supports_tools=True) request = ModelRequest(model=primary, messages=[HumanMessage("Answer")], tools=[]) seen: list[ModelRequest] = [] async def handler(routed: ModelRequest): seen.append(routed) if len(seen) == 1: return AIMessage( content="", additional_kwargs={"reasoning_content": "hidden"} ) return AIMessage(content="Final answer") response = await EvoRouteFallbackMiddleware([]).awrap_model_call(request, handler) assert isinstance(response, AIMessage) assert response.content == "Final answer" assert len(seen) == 2 assert isinstance(seen[1].messages[0], SystemMessage) assert "without final text" in str(seen[1].messages[0].content) @pytest.mark.asyncio async def test_repeated_empty_provider_response_fails_with_specific_code(): primary = _Model("primary", supports_tools=True) attempts = 0 async def handler(_routed: ModelRequest): nonlocal attempts attempts += 1 return AIMessage(content=[{"type": "reasoning", "summary": []}]) with pytest.raises(ModelProviderResponseError) as captured: await EvoRouteFallbackMiddleware([]).awrap_model_call( ModelRequest(model=primary, messages=[], tools=[]), handler ) assert captured.value.code == "MODEL_PROVIDER_RESPONSE_INVALID" assert attempts == 2 @pytest.mark.asyncio async def test_fallback_applies_its_own_tool_capability(): primary = _Model("primary", supports_tools=True) fallback = _Model("fallback", supports_tools=False) tools = [{"type": "function", "function": {"name": "search"}}] request = ModelRequest(model=primary, messages=[], tools=tools) seen = [] async def handler(routed: ModelRequest): seen.append((routed.model.metadata["route_key"], list(routed.tools))) if routed.model is primary: raise ConnectionError("upstream unavailable") return "ok" assert await EvoRouteFallbackMiddleware([fallback]).awrap_model_call( request, handler ) == "ok" assert seen == [("primary", tools), ("fallback", [])] @pytest.mark.asyncio async def test_route_fallback_uses_evo_frozen_fallback_model(): primary = _Model("primary") fallback = _Model("fallback") middleware = EvoRouteFallbackMiddleware([fallback], _Health()) seen: list[str] = [] async def handler(request: ModelRequest): route = request.model.metadata["route_key"] seen.append(route) if route == "primary": raise ConnectionError("upstream unavailable") return "fallback-response" result = await middleware.awrap_model_call(_request(primary), handler) assert result == "fallback-response" assert seen == ["primary", "fallback"] @pytest.mark.asyncio async def test_protocol_error_retries_with_a_bounded_repair_instruction(): primary = _Model("primary", supports_tools=True) request = ModelRequest( model=primary, messages=[], tools=[{"type": "function", "function": {"name": "search"}}], ) seen: list[ModelRequest] = [] async def handler(routed: ModelRequest): seen.append(routed) if len(seen) == 1: raise ModelToolProtocolError("unknown_name") return "repaired-response" result = await EvoRouteFallbackMiddleware([primary]).awrap_model_call( request, handler ) assert result == "repaired-response" assert len(seen) == 2 assert seen[0].messages == [] assert isinstance(seen[1].messages[0], SystemMessage) assert "unknown_name" in str(seen[1].messages[0].content) assert "search" in str(seen[1].messages[0].content) @pytest.mark.asyncio async def test_other_agent_control_errors_remain_non_fallbackable(): primary = _Model("primary") fallback = _Model("fallback") seen: list[str] = [] async def handler(routed: ModelRequest): seen.append(routed.model.metadata["route_key"]) raise AgentControlError("MODEL_REQUEST_REJECTED", "rejected") with pytest.raises(AgentControlError): await EvoRouteFallbackMiddleware([fallback]).awrap_model_call( _request(primary), handler ) assert seen == ["primary"] @pytest.mark.asyncio async def test_open_primary_is_skipped_and_control_errors_do_not_fallback(): primary = _Model("primary") fallback = _Model("fallback") middleware = EvoRouteFallbackMiddleware([fallback], _Health({"primary"})) seen: list[str] = [] async def handler(request: ModelRequest): seen.append(request.model.metadata["route_key"]) return "fallback-response" assert await middleware.awrap_model_call(_request(primary), handler) == "fallback-response" assert seen == ["fallback"] async def controlled(_request: ModelRequest): raise EvoRuntimeError("ADMISSION_EXHAUSTED") with pytest.raises(EvoRuntimeError, match="ADMISSION_EXHAUSTED"): await EvoRouteFallbackMiddleware([fallback]).awrap_model_call( _request(primary), controlled, ) def test_sync_open_primary_is_skipped(): primary = _Model("primary") fallback = _Model("fallback") middleware = EvoRouteFallbackMiddleware([fallback], _Health({"primary"})) seen: list[str] = [] def handler(request: ModelRequest): seen.append(request.model.metadata["route_key"]) return "fallback-response" assert middleware.wrap_model_call(_request(primary), handler) == "fallback-response" assert seen == ["fallback"]