fix multi-query agent workspace searches

This commit is contained in:
Hamza Ayed
2026-10-04 17:42:57 +03:00
parent b3dc6154e5
commit a62a831e90
6 changed files with 253 additions and 17 deletions
+68 -3
View File
@@ -16,6 +16,7 @@ if "SOVEREIGNAI_DATA_DIR" not in os.environ:
from app.main import (
AgentRequest,
_execute_agent,
_explicit_workspace_search_queries,
_knowledge_context_for_model,
_is_knowledge_answer_insufficient,
app,
@@ -46,6 +47,70 @@ class AgentSkillTests(unittest.TestCase):
self.assertIn("propose_file_change", skills["code_review"]["allowed_tools"])
self.assertEqual(response.json()["default"], None)
def test_explicit_numbered_workspace_searches_are_extracted(self) -> None:
task = (
"استخدم أداة البحث في مساحة العمل ثلاث مرات منفصلة للعثور على: "
"(1) AGENT_STREAM_HEARTBEAT_SECONDS، "
"(2) مهلة عميل Flutter لمسار run/stream، "
"(3) timeout_seconds=600.0 في تنفيذ الوكيل. "
"بعد كل بحث افحص النتيجة، ثم لخّص العلاقة بين المهلات."
)
self.assertEqual(
_explicit_workspace_search_queries(task),
[
"AGENT_STREAM_HEARTBEAT_SECONDS",
"مهلة عميل Flutter لمسار run/stream",
"timeout_seconds=600.0 في تنفيذ الوكيل",
],
)
def test_numbered_workspace_searches_are_prefetched_separately(self) -> None:
task = (
"استخدم أداة البحث في مساحة العمل ثلاث مرات منفصلة للعثور على: "
"(1) AGENT_STREAM_HEARTBEAT_SECONDS، "
"(2) مهلة عميل Flutter لمسار run/stream، "
"(3) timeout_seconds=600.0 في تنفيذ الوكيل. "
"بعد كل بحث افحص النتيجة، ثم لخّص العلاقة بين المهلات."
)
queries = _explicit_workspace_search_queries(task)
matches = {
queries[0]: [("app/main.py", "50: AGENT_STREAM_HEARTBEAT_SECONDS = 15.0")],
queries[1]: [("flutter_app/lib/core/network/api_repository.dart", "570: timeout(minutes: 10) run/stream")],
queries[2]: [("app/main.py", "2212: timeout_seconds=600.0")],
}
completion = {"choices": [{"message": {"content": "وجدت القيم الثلاث."}}]}
with (
patch(
"app.main.workspace.retrieve",
side_effect=lambda query, _root, limit=3: matches[query],
) as search,
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
):
result = asyncio.run(
_execute_agent(
AgentRequest(
task=task,
workspace_path=self.workspace,
skill_id="code_explain",
)
)
)
self.assertEqual([call.args[0] for call in search.call_args_list], queries)
self.assertEqual(result["files"], ["app/main.py", "flutter_app/lib/core/network/api_repository.dart"])
self.assertEqual(result["steps"], [{"tool": "search_workspace", "status": "completed"}])
evidence = model.await_args.args[0]["messages"][1]["content"]
system = model.await_args.args[0]["messages"][0]["content"]
offered_tools = [
item["function"]["name"] for item in model.await_args.args[0]["tools"]
]
self.assertIn("AGENT_STREAM_HEARTBEAT_SECONDS = 15.0", evidence)
self.assertIn("timeout(minutes: 10)", evidence)
self.assertIn("timeout_seconds=600.0", evidence)
self.assertIn("لا للحلقة كاملة", system)
self.assertIn("search_workspace", offered_tools)
def test_safe_arithmetic_accepts_common_unicode_operator_symbols(self) -> None:
self.assertEqual(safe_arithmetic("137 × 29"), 3973.0)
self.assertEqual(safe_arithmetic("12 ÷ 3 − 1"), 3.0)
@@ -451,7 +516,7 @@ class AgentSkillTests(unittest.TestCase):
)
)
search.assert_called_once_with(task, Path(self.workspace).resolve())
search.assert_called_once_with(task, Path(self.workspace).resolve(), limit=3)
initial_payload = model.await_args_list[0].args[0]
self.assertNotIn(
"search_workspace",
@@ -495,7 +560,7 @@ class AgentSkillTests(unittest.TestCase):
)
)
search.assert_called_once_with(task, Path(self.workspace).resolve())
search.assert_called_once_with(task, Path(self.workspace).resolve(), limit=3)
model.assert_not_awaited()
self.assertEqual(
[step["tool"] for step in result["steps"]],
@@ -527,7 +592,7 @@ class AgentSkillTests(unittest.TestCase):
)
self.assertEqual(error.exception.status_code, 422)
search.assert_called_once_with(task, Path(self.workspace).resolve())
search.assert_called_once_with(task, Path(self.workspace).resolve(), limit=3)
offered = model.await_args.args[0]["tools"]
self.assertNotIn("search_workspace", [item["function"]["name"] for item in offered])
@@ -111,7 +111,7 @@ class TimeoutAndCancellationTests(unittest.IsolatedAsyncioTestCase):
await stream.aclose()
await asyncio.wait_for(cancelled.wait(), timeout=1)
async def test_idle_agent_stream_emits_keepalive_until_model_finishes(self) -> None:
async def test_idle_agent_stream_emits_heartbeat_until_model_finishes(self) -> None:
started = asyncio.Event()
cancelled = asyncio.Event()
@@ -131,8 +131,8 @@ class TimeoutAndCancellationTests(unittest.IsolatedAsyncioTestCase):
):
response = await run_agent_stream(request, user_id=database.LOCAL_USER_ID)
stream = response.body_iterator
keepalive = await asyncio.wait_for(anext(stream), timeout=1)
self.assertEqual(keepalive, ": keep-alive\n\n")
heartbeat = await asyncio.wait_for(anext(stream), timeout=1)
self.assertEqual(heartbeat, "event: heartbeat\ndata: {}\n\n")
self.assertTrue(started.is_set())
await stream.aclose()
await asyncio.wait_for(cancelled.wait(), timeout=1)
@@ -0,0 +1,32 @@
import unittest
from pathlib import Path
from app import workspace
class WorkspaceSearchTests(unittest.TestCase):
def test_search_excerpt_returns_numbered_matching_lines_from_deep_in_file(self) -> None:
project_root = Path(__file__).resolve().parents[1]
results = workspace.retrieve("timeout_seconds=600.0", project_root, limit=10)
path, excerpt = next(
(path, excerpt)
for path, excerpt in results
if path == "app/main.py"
)
self.assertRegex(excerpt, r"\d+: .*timeout_seconds=600\.0")
self.assertNotIn("from __future__ import annotations", excerpt)
def test_search_ranks_requested_flutter_timeout_code_before_docs(self) -> None:
project_root = Path(__file__).resolve().parents[1]
results = workspace.retrieve(
"مهلة عميل Flutter لمسار run/stream", project_root, limit=3
)
path, excerpt = results[0]
self.assertEqual(path, "flutter_app/lib/core/network/api_repository.dart")
self.assertIn("minutes: 10", excerpt)
if __name__ == "__main__":
unittest.main()