"""Final model tool protocol validation tests.""" from dataclasses import dataclass, field, replace from types import SimpleNamespace from typing import Any import pytest from langchain.agents.middleware.types import ExtendedModelResponse, ModelResponse from langchain_core.messages import AIMessage from EvoScientist.llm.errors import ModelToolProtocolError from EvoScientist.middleware.tool_protocol_guard import ToolProtocolGuardMiddleware @dataclass(frozen=True) class _Request: tools: list[Any] model: Any = field(default_factory=lambda: SimpleNamespace(metadata={})) def override(self, **updates: Any): return replace(self, **updates) def _response(*calls: dict[str, Any], content: Any = "") -> ModelResponse: return ModelResponse(result=[AIMessage(content=content, tool_calls=list(calls))]) def _call(call_id: str = "call-1", name: str = "search", args: Any = None): return {"id": call_id, "name": name, "args": {} if args is None else args} @pytest.mark.parametrize( ("call", "reason"), [ (_call(name=""), "missing_name"), (_call(name=" "), "missing_name"), (_call(name="missing"), "unknown_name"), ], ) def test_invalid_final_tool_call_fails_closed(call, reason): middleware = ToolProtocolGuardMiddleware() request = _Request(tools=[{"name": "search"}]) with pytest.raises(ModelToolProtocolError) as caught: middleware.wrap_model_call(request, lambda _request: _response(call)) assert caught.value.reason == reason assert caught.value.retryable is True assert caught.value.fallbackable is True assert caught.value.non_fallbackable is False def test_json_string_args_are_normalized_to_an_object(): message = AIMessage(content="", tool_calls=[_call()]) message.tool_calls[0]["args"] = "{}" result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) assert result.result[0].tool_calls[0]["args"] == {} assert message.tool_calls[0]["args"] == "{}" def test_malformed_json_args_are_rejected(): message = AIMessage(content="", tool_calls=[_call()]) message.tool_calls[0]["args"] = "{" with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) assert caught.value.reason == "invalid_args" def test_duplicate_parallel_call_id_rejects_whole_response(): request = _Request(tools=[{"name": "search"}, {"name": "read_file"}]) response = _response(_call(name="search"), _call(name="read_file")) with pytest.raises( ModelToolProtocolError, match="invalid structured tool call" ) as caught: ToolProtocolGuardMiddleware().wrap_model_call( request, lambda _request: response ) assert caught.value.reason == "duplicate_id" def test_one_invalid_parallel_call_rejects_atomically(): request = _Request(tools=[{"name": "search"}]) response = _response(_call(call_id="one"), _call(call_id="two", name="missing")) with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( request, lambda _request: response ) assert caught.value.reason == "unknown_name" def test_final_invalid_tool_calls_are_rejected(): message = AIMessage( content="", invalid_tool_calls=[ {"id": "bad", "name": "search", "args": "{", "error": "bad json"} ], ) with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) assert caught.value.reason == "invalid_final_call" assert caught.value.call_id == "bad" def test_content_block_must_match_parsed_call(): response = _response( _call(), content=[ {"type": "tool_call", "id": "call-1", "name": "read_file", "args": {}} ], ) with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}, {"name": "read_file"}]), lambda _request: response, ) assert caught.value.reason == "inconsistent_source" def test_responses_content_block_uses_call_id_over_output_item_id(): """Responses item IDs are not the identifiers used for tool results.""" response = _response( _call(call_id="call-result-1", name="search", args={}), content=[ { "type": "function_call", "id": "fc-output-item-1", "call_id": "call-result-1", "name": "search", "arguments": "{}", } ], ) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: response ) assert result.result[0].tool_calls[0]["id"] == "call-result-1" def test_parsed_only_valid_call_and_extended_response_pass(): response = ExtendedModelResponse(model_response=_response(_call())) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"type": "function", "function": {"name": "search"}}]), lambda _request: response, ) assert result is response async def test_async_direct_ai_message_shape_passes(): response = AIMessage(content="", tool_calls=[_call()]) async def handler(_request): return response result = await ToolProtocolGuardMiddleware().awrap_model_call( _Request(tools=[{"name": "search"}]), handler ) assert result is response def test_error_carries_safe_route_metadata(): model = SimpleNamespace( metadata={ "route_provider": "openai", "route_model": "gpt-example", "route_key": "route-safe", "route_config_generation": 12, "route_api_mode": "chat_completions", "route_endpoint": "primary", "route_tool_call_transport": "streaming", } ) with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}], model=model), lambda _request: _response(_call(name="")), ) payload = caught.value.model_dump() assert payload["route_key"] == "route-safe" assert payload["config_generation"] == 12 assert payload["endpoint"] == "primary" assert payload["tool_call_transport"] == "streaming" assert "args" not in payload def test_missing_id_is_generated_without_mutating_the_provider_message(): call = _call(call_id="", name="search", args={"query": "private search text"}) provider_response = _response(call) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: provider_response, ) normalized = result.result[0].tool_calls[0] assert normalized["id"].startswith("call_") assert normalized["args"] == {"query": "private search text"} assert provider_response.result[0].tool_calls[0]["id"] == "" def test_generic_openai_compatible_parallel_calls_without_ids_are_stable(): model = SimpleNamespace( metadata={ "route_adapter_id": "generic-openai-compatible", "route_key": "qwen-primary", "route_config_generation": 7, "route_api_mode": "chat_completions", } ) response = _response( _call(call_id="", name="think_tool", args={"reflection": "plan"}), _call(call_id="", name="execute", args={"command": "pwd"}), ) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request( tools=[{"name": "think_tool"}, {"name": "execute"}], model=model ), lambda _request: response, ) calls = result.result[0].tool_calls assert [call["name"] for call in calls] == ["think_tool", "execute"] assert all(call["id"].startswith("call_") for call in calls) assert calls[0]["id"] != calls[1]["id"] assert response.result[0].tool_calls[0]["id"] == "" def test_raw_openai_id_is_merged_into_the_canonical_call(): parsed = _call(call_id="", name="search", args={"query": "secret"}) raw = { "id": "provider-call-id", "type": "function", "function": {"name": "search", "arguments": '{"query":"secret"}'}, } message = AIMessage( content="", tool_calls=[parsed], additional_kwargs={"tool_calls": [raw]}, ) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) normalized = result.result[0] assert normalized.tool_calls == [ { "id": "provider-call-id", "name": "search", "args": {"query": "secret"}, "type": "tool_call", } ] assert "tool_calls" not in normalized.additional_kwargs assert message.additional_kwargs["tool_calls"] == [raw] def test_content_only_function_call_is_decoded_and_normalized(): message = AIMessage( content=[ { "type": "function_call", "id": "", "name": "search", "arguments": '{"query":"x"}', } ] ) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) normalized = result.result[0] call_id = normalized.tool_calls[0]["id"] assert call_id.startswith("call_") assert normalized.tool_calls[0]["args"] == {"query": "x"} assert normalized.content[0]["id"] == call_id def test_legacy_function_call_is_decoded_and_removed_from_replay_metadata(): legacy = {"name": "search", "arguments": '{"query":"x"}'} message = AIMessage(content="", additional_kwargs={"function_call": legacy}) result = ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) normalized = result.result[0] assert normalized.tool_calls[0]["id"].startswith("call_") assert normalized.tool_calls[0]["args"] == {"query": "x"} assert "function_call" not in normalized.additional_kwargs assert message.additional_kwargs["function_call"] == legacy def test_diagnostic_failure_cannot_mask_the_protocol_error(): circular: dict[str, Any] = {} circular["self"] = circular message = AIMessage(content="", tool_calls=[_call(call_id="", args={})]) message.tool_calls[0]["args"] = circular with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) assert caught.value.reason == "missing_id" assert caught.value.call_diagnostic["args_digest"].startswith("sha256:") def test_protocol_failure_logs_only_redacted_call_diagnostic(caplog): call = _call(name="", args={"query": "private search text"}) with caplog.at_level("WARNING"), pytest.raises(ModelToolProtocolError): ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: _response(call), ) record = caplog.records[-1].getMessage() assert "reason=missing_name" in record assert '"args_keys": ["query"]' in record assert "private search text" not in record