Improve bounded agent tool orchestration

This commit is contained in:
Hamza Ayed
2026-10-03 16:46:27 +03:00
parent 65812cf3ee
commit ac8359de6f
5 changed files with 423 additions and 152 deletions
@@ -1,4 +1,5 @@
import asyncio
import json
import os
import tempfile
import unittest
@@ -97,6 +98,61 @@ class AgentSkillTests(unittest.TestCase):
self.assertEqual(result["skill"], "code_review")
self.assertNotIn("proposal", result)
def test_model_cannot_call_a_tool_after_file_proposal_ends_tool_access(self) -> None:
completions = [
{"choices": [{"message": {"tool_calls": [{
"id": "proposal-1",
"function": {
"name": "propose_file_change",
"arguments": '{"path":"new.py","operation":"create","content":"print(1)"}',
},
}]}}]},
{"choices": [{"message": {"tool_calls": [{
"id": "late-calc",
"function": {"name": "calculator", "arguments": '{"expression":"1 + 1"}'},
}]}}]},
]
proposal = {
"path": "new.py",
"operation": "create",
"expires_in_seconds": 600,
}
with (
patch("app.main.workspace.create_change_preview", return_value=proposal) as create_preview,
patch("app.main.get_completion", new=AsyncMock(side_effect=completions)) as model,
):
with self.assertRaises(HTTPException) as error:
asyncio.run(
_execute_agent(
AgentRequest(
task="أنشئ معاينة ملف new.py",
workspace_path=self.workspace,
skill_id="code_review",
)
)
)
self.assertEqual(error.exception.status_code, 422)
self.assertEqual(create_preview.call_count, 1)
self.assertEqual(model.await_count, 2)
def test_model_tool_names_must_be_strings_from_the_offered_tool_set(self) -> None:
malformed = {
"choices": [{"message": {"tool_calls": [{
"id": "bad-name",
"function": {"name": ["calculator"], "arguments": "{}"},
}]}}]
}
with patch("app.main.get_completion", new=AsyncMock(return_value=malformed)):
with self.assertRaises(HTTPException) as error:
asyncio.run(
_execute_agent(
AgentRequest(task="احسب 1+1", skill_id="code_explain")
)
)
self.assertEqual(error.exception.status_code, 422)
def test_preselected_file_is_read_once_without_redundant_search_call(self) -> None:
completion = {"choices": [{"message": {"content": "المهارات مسجلة في قاموس محلي."}}]}
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
@@ -194,6 +250,109 @@ class AgentSkillTests(unittest.TestCase):
self.assertEqual(error.exception.status_code, 422)
def test_agent_can_chain_read_only_search_and_calculation(self) -> None:
def tool_call(call_id: str, name: str, arguments: dict[str, str]) -> dict:
return {
"id": call_id,
"type": "function",
"function": {"name": name, "arguments": json.dumps(arguments)},
}
completions = [
{"choices": [{"message": {"tool_calls": [
tool_call("search-1", "search_workspace", {"query": "secret key config"})
]}}]},
{"choices": [{"message": {"tool_calls": [
tool_call("calc-1", "calculator", {"expression": "19 * 23"})
]}}]},
{"choices": [{"message": {"content": "وجدت الإعداد، والحساب يساوي 437."}}]},
]
with (
patch("app.main.workspace.retrieve", return_value=[("app/config.py", "key comes from env")]) as search,
patch("app.main.get_completion", new=AsyncMock(side_effect=completions)) as model,
):
result = asyncio.run(
_execute_agent(
AgentRequest(
task="ابحث عن مصدر المفتاح واحسب 19 في 23",
workspace_path=self.workspace,
skill_id="code_explain",
)
)
)
search.assert_called_once()
self.assertEqual(model.await_count, 3)
self.assertEqual(
[step["tool"] for step in result["steps"]],
["search_workspace", "calculator"],
)
self.assertEqual(result["files"], ["app/config.py"])
self.assertIn("437", result["result"])
second_messages = model.await_args_list[1].args[0]["messages"]
self.assertEqual(second_messages[-1]["role"], "tool")
self.assertIn("key comes from env", second_messages[-1]["content"])
third_messages = model.await_args_list[2].args[0]["messages"]
self.assertEqual([message["role"] for message in third_messages[-2:]], ["assistant", "tool"])
def test_explicit_workspace_search_prefetches_before_followup_tool(self) -> None:
task = "ابحث في ملفات المشروع عن قيمة MAX_AGENT_TOOL_CALLS، ثم احسب القيمة مضروبة في 7."
completions = [
{"choices": [{"message": {"tool_calls": [{
"id": "calc-after-search",
"function": {"name": "calculator", "arguments": '{"expression":"3 * 7"}'},
}]}}]},
{"choices": [{"message": {"content": "القيمة 3، والناتج 21 من app/main.py."}}]},
]
with (
patch("app.main.workspace.retrieve", return_value=[("app/main.py", "MAX_AGENT_TOOL_CALLS = 3")]) as search,
patch("app.main.get_completion", new=AsyncMock(side_effect=completions)) as model,
):
result = asyncio.run(
_execute_agent(
AgentRequest(
task=task,
workspace_path=self.workspace,
skill_id="code_explain",
)
)
)
search.assert_called_once_with(task, Path(self.workspace).resolve())
initial_payload = model.await_args_list[0].args[0]
self.assertIn(
"search_workspace",
[tool["function"]["name"] for tool in initial_payload["tools"]],
)
self.assertIn("MAX_AGENT_TOOL_CALLS = 3", initial_payload["messages"][1]["content"])
self.assertEqual(
[step["tool"] for step in result["steps"]],
["search_workspace", "calculator"],
)
self.assertEqual(result["files"], ["app/main.py"])
self.assertIn("21", result["result"])
def test_agent_caps_sequential_tools_and_forces_final_model_turn(self) -> None:
tool_calls = [
{"choices": [{"message": {"tool_calls": [{
"id": f"calc-{index}",
"function": {"name": "calculator", "arguments": '{"expression":"2 + 3"}'},
}]}}]}
for index in range(3)
]
tool_calls.append({"choices": [{"message": {"content": "انتهيت بعد ثلاث خطوات."}}]})
with patch("app.main.get_completion", new=AsyncMock(side_effect=tool_calls)) as model:
result = asyncio.run(
_execute_agent(
AgentRequest(task="استخدم الحاسبة ثلاث مرات", skill_id="code_explain")
)
)
self.assertEqual(model.await_count, 4)
self.assertEqual(len(result["steps"]), 3)
self.assertNotIn("tools", model.await_args_list[3].args[0])
self.assertEqual(result["result"], "انتهيت بعد ثلاث خطوات.")
def test_unknown_skill_is_rejected_by_request_contract(self) -> None:
response = self.client.post(
"/v1/agent/run",