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
Introduce provider, model, and invocation contracts with encrypted configuration persistence. Add web runtime fencing, route fallback, recovery middleware, workspace scoping, and comprehensive tests.
354 lines
12 KiB
Python
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
|