"""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"), (_call(call_id=""), "missing_id"), ], ) 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 False assert caught.value.fallbackable is True def test_non_mapping_args_are_rejected_if_adapter_bypasses_message_validation(): 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_block" 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_carries_redacted_call_diagnostic_only_for_internal_logging(): call = _call(call_id="", name="search", args={"query": "private search text"}) with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: _response(call), ) diagnostic = caught.value.call_diagnostic assert diagnostic == { "source": "parsed_tool_calls", "call_index": 0, "call_count": 1, "call_type": "object", "name": "search", "id_present": False, "args_present": True, "args_type": "object", "args_key_count": 1, "args_keys": ["query"], "args_keys_truncated": False, "args_digest": diagnostic["args_digest"], "raw_openai_call_available": False, } assert diagnostic["args_digest"].startswith("sha256:") assert "private search text" not in str(diagnostic) assert "call_diagnostic" not in caught.value.model_dump() def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes(): 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]}, ) with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), lambda _request: ModelResponse(result=[message]), ) diagnostic = caught.value.call_diagnostic assert diagnostic["id_present"] is False assert diagnostic["raw_openai_call_available"] is True assert diagnostic["raw_openai_call"]["id_present"] is True assert diagnostic["raw_openai_call"]["name"] == "search" assert "provider-call-id" not in str(diagnostic) assert "secret" not in str(diagnostic) 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:")