mirror of
https://hubproxy.babadafafafafa.cn/https://github.com/usestrix/strix.git
synced 2026-09-20 08:03:42 +08:00
fix(agents): tell the model to re-issue tool calls whose arguments are not valid JSON
A streamed turn that ends mid-tool-call leaves the client with a JSON prefix as the arguments. Instead of the SDK's generic parse-error result, the tool wrapper now short-circuits with a clear 'not executed, re-issue with complete arguments' message so the model recovers in one turn.
This commit is contained in:
@@ -18,6 +18,7 @@ from pydantic import ValidationError
|
|||||||
|
|
||||||
from strix.agents.prompt import render_system_prompt
|
from strix.agents.prompt import render_system_prompt
|
||||||
from strix.config import load_settings
|
from strix.config import load_settings
|
||||||
|
from strix.config.tool_call_arguments import describe_malformed_arguments
|
||||||
from strix.tools.agents_graph.tools import (
|
from strix.tools.agents_graph.tools import (
|
||||||
agent_finish,
|
agent_finish,
|
||||||
create_agent,
|
create_agent,
|
||||||
@@ -260,6 +261,10 @@ def _with_coerced_arguments(tool: FunctionTool) -> FunctionTool:
|
|||||||
nullish = tool.name.startswith(_QUERY_TOOL_PREFIXES)
|
nullish = tool.name.startswith(_QUERY_TOOL_PREFIXES)
|
||||||
|
|
||||||
async def invoke(ctx: Any, raw_input: str) -> Any:
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
|
malformed = describe_malformed_arguments(tool.name, raw_input)
|
||||||
|
if malformed is not None:
|
||||||
|
logger.debug("Tool %s got malformed arguments; asking the model to re-issue", tool.name)
|
||||||
|
return malformed
|
||||||
return await invoke_tool(ctx, _coerce_arguments(raw_input, schema, nullish=nullish))
|
return await invoke_tool(ctx, _coerce_arguments(raw_input, schema, nullish=nullish))
|
||||||
|
|
||||||
tool.on_invoke_tool = invoke
|
tool.on_invoke_tool = invoke
|
||||||
|
|||||||
@@ -22,6 +22,31 @@ from openai.types.responses import ResponseFunctionToolCall
|
|||||||
MALFORMED_ARGUMENTS_KEY = "malformed_arguments"
|
MALFORMED_ARGUMENTS_KEY = "malformed_arguments"
|
||||||
|
|
||||||
|
|
||||||
|
def describe_malformed_arguments(tool_name: str, arguments: str) -> str | None:
|
||||||
|
"""Return a model-facing recovery message if ``arguments`` is not a JSON object.
|
||||||
|
|
||||||
|
A tool call whose arguments do not parse is almost always one the stream cut
|
||||||
|
off (the server flushed a prefix of the JSON, or the client closed early),
|
||||||
|
so instead of the SDK's generic parse-error result the model is told the
|
||||||
|
call never ran and must be re-issued whole.
|
||||||
|
"""
|
||||||
|
if not arguments.strip():
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
parsed = json.loads(arguments)
|
||||||
|
except ValueError as exc:
|
||||||
|
detail = str(exc)
|
||||||
|
else:
|
||||||
|
if isinstance(parsed, dict):
|
||||||
|
return None
|
||||||
|
detail = f"expected a JSON object, got {type(parsed).__name__}"
|
||||||
|
return (
|
||||||
|
f"{tool_name}: the tool call was not executed because its arguments were "
|
||||||
|
f"truncated or otherwise not valid JSON ({detail}). The response was likely "
|
||||||
|
"cut off mid-call. Re-issue the call with complete, valid JSON arguments."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def repair_arguments(arguments: object) -> str | None:
|
def repair_arguments(arguments: object) -> str | None:
|
||||||
"""Return replacement arguments a strict server accepts, or ``None`` if already valid."""
|
"""Return replacement arguments a strict server accepts, or ``None`` if already valid."""
|
||||||
if not isinstance(arguments, str) or not arguments.strip():
|
if not isinstance(arguments, str) or not arguments.strip():
|
||||||
|
|||||||
@@ -38,6 +38,20 @@ async def _roundtrip(
|
|||||||
|
|
||||||
|
|
||||||
_STRING = {"todos": {"type": "string"}}
|
_STRING = {"todos": {"type": "string"}}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_truncated_arguments_short_circuit_with_a_reissue_message() -> None:
|
||||||
|
captured: dict[str, str] = {}
|
||||||
|
wrapped = factory._with_coerced_arguments(_capturing_tool(captured, _STRING))
|
||||||
|
|
||||||
|
result = await wrapped.on_invoke_tool(cast("Any", None), '{"todos": "a, b')
|
||||||
|
|
||||||
|
assert "not executed" in result
|
||||||
|
assert "Re-issue the call" in result
|
||||||
|
assert "raw_input" not in captured
|
||||||
|
|
||||||
|
|
||||||
_ARRAY = {"tags": {"type": "array", "items": {"type": "string"}}}
|
_ARRAY = {"tags": {"type": "array", "items": {"type": "string"}}}
|
||||||
_NULLABLE_ARRAY = {
|
_NULLABLE_ARRAY = {
|
||||||
"tags": {"anyOf": [{"type": "array", "items": {"type": "string"}}, {"type": "null"}]}
|
"tags": {"anyOf": [{"type": "array", "items": {"type": "string"}}, {"type": "null"}]}
|
||||||
@@ -136,12 +150,14 @@ async def test_unknown_and_null_arguments_are_untouched() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_non_object_payloads_pass_through_unchanged() -> None:
|
async def test_non_object_payloads_are_reported_instead_of_invoked() -> None:
|
||||||
captured: dict[str, str] = {}
|
captured: dict[str, str] = {}
|
||||||
wrapped = factory._with_coerced_arguments(_capturing_tool(captured, _ARRAY))
|
wrapped = factory._with_coerced_arguments(_capturing_tool(captured, _ARRAY))
|
||||||
|
|
||||||
assert await wrapped.on_invoke_tool(cast("Any", None), "not json") == "ok"
|
result = await wrapped.on_invoke_tool(cast("Any", None), "not json")
|
||||||
assert captured["raw_input"] == "not json"
|
|
||||||
|
assert result.startswith("probe: the tool call was not executed")
|
||||||
|
assert "raw_input" not in captured
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from openai.types.responses import ResponseFunctionToolCall
|
|||||||
from strix.config.models import _NonStreamingModel, _TurnGuardModel
|
from strix.config.models import _NonStreamingModel, _TurnGuardModel
|
||||||
from strix.config.tool_call_arguments import (
|
from strix.config.tool_call_arguments import (
|
||||||
MALFORMED_ARGUMENTS_KEY,
|
MALFORMED_ARGUMENTS_KEY,
|
||||||
|
describe_malformed_arguments,
|
||||||
repair_arguments,
|
repair_arguments,
|
||||||
repair_history_arguments,
|
repair_history_arguments,
|
||||||
repair_input,
|
repair_input,
|
||||||
@@ -244,3 +245,17 @@ def test_repair_input_returns_same_object_when_nothing_changes() -> None:
|
|||||||
|
|
||||||
assert repair_input(items) is items
|
assert repair_input(items) is items
|
||||||
assert repair_input("plain prompt") == "plain prompt"
|
assert repair_input("plain prompt") == "plain prompt"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("arguments", ['{"cmd": "ls -la', "[1, 2]", "null", "not json"])
|
||||||
|
def test_describe_malformed_arguments_tells_the_model_to_reissue(arguments: str) -> None:
|
||||||
|
message = describe_malformed_arguments("exec_command", arguments)
|
||||||
|
|
||||||
|
assert message is not None
|
||||||
|
assert message.startswith("exec_command: the tool call was not executed")
|
||||||
|
assert "Re-issue the call" in message
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("arguments", ["", " ", "{}", '{"cmd": "ls"}'])
|
||||||
|
def test_describe_malformed_arguments_accepts_objects_and_empty_input(arguments: str) -> None:
|
||||||
|
assert describe_malformed_arguments("exec_command", arguments) is None
|
||||||
|
|||||||
Reference in New Issue
Block a user