Files
EvoScientist-Multi/tests/test_tool_protocol_guard.py
T

247 lines
8.2 KiB
Python

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