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()