136 lines
5.6 KiB
Python
136 lines
5.6 KiB
Python
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from owner_voice_pet.agent_memory import FakeMemoryManager, MemoryRecordInput
|
|
from owner_voice_pet.tool_router import (
|
|
FakeToolAdapter,
|
|
MemorySaveTool,
|
|
MemorySearchTool,
|
|
ShellReadonlyTool,
|
|
ToolCallRequest,
|
|
ToolContext,
|
|
ToolRouter,
|
|
ToolRiskClassifier,
|
|
)
|
|
|
|
|
|
class ToolRouterTests(unittest.TestCase):
|
|
def test_unknown_tool_is_rejected(self) -> None:
|
|
router = ToolRouter({})
|
|
request = ToolCallRequest("1", "unknown.tool", {}, "turn-1")
|
|
|
|
decision = router.route(request, ToolContext())
|
|
|
|
self.assertEqual(decision.action, "reject")
|
|
self.assertEqual(decision.risk_level, "forbidden")
|
|
|
|
def test_memory_search_tool_executes_and_returns_sanitized_result(self) -> None:
|
|
memory = FakeMemoryManager()
|
|
memory.save(MemoryRecordInput("preference", "用户喜欢 Python"))
|
|
router = ToolRouter({"memory.search": MemorySearchTool()})
|
|
request = ToolCallRequest("1", "memory.search", {"query": "Python", "top_k": 1}, "turn-1")
|
|
|
|
decision = router.route(request, ToolContext(memory=memory))
|
|
result = router.execute(request, decision, ToolContext(memory=memory))
|
|
|
|
self.assertEqual(decision.action, "execute")
|
|
self.assertEqual(result.status, "success")
|
|
self.assertIn("用户喜欢 Python", result.output_text)
|
|
self.assertEqual(router.audit_log[-1].tool_name, "memory.search")
|
|
|
|
def test_memory_save_sensitive_text_requires_confirmation(self) -> None:
|
|
router = ToolRouter({"memory.save": MemorySaveTool()})
|
|
request = ToolCallRequest(
|
|
"1",
|
|
"memory.save",
|
|
{"text": "保存 api key secret-value"},
|
|
"turn-1",
|
|
)
|
|
|
|
decision = router.route(request, ToolContext(memory=FakeMemoryManager()))
|
|
result = router.execute(request, decision, ToolContext(memory=FakeMemoryManager()))
|
|
|
|
self.assertEqual(decision.action, "require_confirmation")
|
|
self.assertEqual(result.status, "confirmation_required")
|
|
self.assertEqual(result.error_code, "TOOL_CONFIRMATION_REQUIRED")
|
|
|
|
def test_shell_readonly_rejects_mutating_command(self) -> None:
|
|
router = ToolRouter({"shell.readonly": ShellReadonlyTool()})
|
|
request = ToolCallRequest("1", "shell.readonly", {"command": "rm -rf /tmp/x"}, "turn-1")
|
|
|
|
decision = router.route(request, ToolContext())
|
|
|
|
self.assertEqual(decision.action, "reject")
|
|
self.assertEqual(decision.risk_level, "forbidden")
|
|
|
|
def test_shell_readonly_executes_allowed_command(self) -> None:
|
|
router = ToolRouter({"shell.readonly": ShellReadonlyTool()})
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
request = ToolCallRequest("1", "shell.readonly", {"command": "pwd"}, "turn-1")
|
|
decision = router.route(request, ToolContext(cwd=Path(tmp)))
|
|
result = router.execute(request, decision, ToolContext(cwd=Path(tmp)))
|
|
|
|
self.assertEqual(decision.action, "execute")
|
|
self.assertEqual(result.status, "success")
|
|
self.assertIn(tmp, result.output_text)
|
|
|
|
def test_tool_budget_and_duplicate_calls_are_rejected(self) -> None:
|
|
adapter = FakeToolAdapter("memory.search", output="ok")
|
|
router = ToolRouter({"memory.search": adapter}, max_calls_per_turn=2)
|
|
request = ToolCallRequest("1", "memory.search", {"query": "a"}, "turn-1")
|
|
decision = router.route(request, ToolContext())
|
|
router.execute(request, decision, ToolContext())
|
|
|
|
duplicate = router.route(request, ToolContext())
|
|
second = ToolCallRequest("2", "memory.search", {"query": "b"}, "turn-1")
|
|
second_decision = router.route(second, ToolContext())
|
|
router.execute(second, second_decision, ToolContext())
|
|
third = ToolCallRequest("3", "memory.search", {"query": "c"}, "turn-1")
|
|
over_budget = router.route(third, ToolContext())
|
|
|
|
self.assertEqual(duplicate.action, "reject")
|
|
self.assertEqual(duplicate.reason, "duplicate tool call rejected")
|
|
self.assertEqual(over_budget.action, "reject")
|
|
self.assertEqual(over_budget.reason, "tool call budget exceeded")
|
|
|
|
def test_output_is_redacted_and_truncated(self) -> None:
|
|
adapter = FakeToolAdapter("memory.search", output="tp-" + "abcdefghijklmnop " + "x" * 100)
|
|
router = ToolRouter({"memory.search": adapter}, output_limit=20)
|
|
request = ToolCallRequest("1", "memory.search", {"query": "secret"}, "turn-1")
|
|
|
|
decision = router.route(request, ToolContext())
|
|
result = router.execute(request, decision, ToolContext())
|
|
|
|
self.assertTrue(result.output_truncated)
|
|
self.assertIn("[redacted]", result.output_text)
|
|
self.assertNotIn("tp-" + "abcdefghijklmnop", result.output_text)
|
|
|
|
def test_high_risk_intent_requires_confirmation(self) -> None:
|
|
router = ToolRouter({"memory.search": FakeToolAdapter("memory.search")})
|
|
request = ToolCallRequest(
|
|
"1",
|
|
"memory.search",
|
|
{"query": "账号"},
|
|
"turn-1",
|
|
natural_language_intent="上传账号资料",
|
|
)
|
|
|
|
decision = router.route(request, ToolContext())
|
|
|
|
self.assertEqual(decision.action, "require_confirmation")
|
|
self.assertEqual(decision.risk_level, "high")
|
|
|
|
def test_risk_classifier_marks_readonly_as_low(self) -> None:
|
|
risk = ToolRiskClassifier().classify(
|
|
ToolCallRequest("1", "shell.readonly", {"command": "git status --short"}, "turn-1")
|
|
)
|
|
|
|
self.assertEqual(risk, "low")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|