fix: harden tool-call protocol and fallback handling

This commit is contained in:
m4
2026-07-19 12:05:56 +08:00
parent 4fc74e7da7
commit 3ce5614254
23 changed files with 2382 additions and 117 deletions
+246
View File
@@ -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:")