fix: harden tool-call protocol and fallback handling
This commit is contained in:
@@ -0,0 +1,246 @@
|
||||
"""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:")
|
||||
Reference in New Issue
Block a user