Improve bounded agent tool orchestration
This commit is contained in:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user