Files
EvoScientist-Multi/tests/test_tool_protocol_guard.py
m4 5a581c78a2
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
feat: add scoped model runtime configuration
Introduce provider, model, and invocation contracts with encrypted configuration persistence. Add web runtime fencing, route fallback, recovery middleware, workspace scoping, and comprehensive tests.
2026-08-14 22:03:04 +08:00

354 lines
12 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"),
],
)
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