From af7eb25ba643f36dfc31f1380a48886cb3284744 Mon Sep 17 00:00:00 2001 From: mkbk Date: Thu, 18 Jun 2026 22:08:24 +0800 Subject: [PATCH] =?UTF-8?q?[Tool=20Router]=EF=BC=9A=E5=AE=8C=E6=88=90?= =?UTF-8?q?=E5=AE=89=E5=85=A8=E5=B7=A5=E5=85=B7=E8=B7=AF=E7=94=B1=E9=AA=A8?= =?UTF-8?q?=E6=9E=B6=EF=BC=8C=E5=8C=85=E5=90=AB=E7=BB=93=E6=9E=84=E5=8C=96?= =?UTF-8?q?=E8=B0=83=E7=94=A8=E3=80=81=E9=A3=8E=E9=99=A9=E5=88=86=E7=B1=BB?= =?UTF-8?q?=E5=92=8C=E5=8F=AA=E8=AF=BB=E5=B7=A5=E5=85=B7=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../tasks.md | 16 +- src/owner_voice_pet/__init__.py | 22 ++ src/owner_voice_pet/tool_router.py | 314 ++++++++++++++++++ tests/test_agent_memory.py | 2 +- tests/test_tool_router.py | 135 ++++++++ 5 files changed, 480 insertions(+), 9 deletions(-) create mode 100644 src/owner_voice_pet/tool_router.py create mode 100644 tests/test_tool_router.py diff --git a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md index 6eaee5c..ecbeb97 100644 --- a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md +++ b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md @@ -59,14 +59,14 @@ ## 7. Tool Router 与安全工具 -- [ ] 7.1 定义 ToolCallRequest schema;前置条件:LLM tool call contract 完成;优先级:P0;验收标准:包含 name、arguments、timeout、turn id、intent;测试要点:未知字段拒绝。 -- [ ] 7.2 定义 ToolDecision 和 ToolResult;前置条件:ToolCallRequest 完成;优先级:P0;验收标准:支持 execute、reject、require_confirmation、cancelled;测试要点:结果脱敏和截断。 -- [ ] 7.3 实现 Tool Router fake core;前置条件:schema 完成;优先级:P0;验收标准:按工具名路由 fake adapters;测试要点:合法、非法、超时、超预算均覆盖。 -- [ ] 7.4 实现工具风险分类规则;前置条件:安全策略确认;优先级:P0;验收标准:低/中/高/禁止风险可解释;测试要点:删除、上传、账号、交易、安装依赖触发高风险或禁止。 -- [ ] 7.5 规划 `memory.search` 和 `memory.save`;前置条件:MemoryManager fake 完成;优先级:P0;验收标准:search 只读,save 走敏感策略;测试要点:敏感 save 需要确认或拒绝。 -- [ ] 7.6 规划 `shell.readonly` adapter;前置条件:allowlist 策略确认;优先级:P0;验收标准:只允许只读命令和受限目录;测试要点:写文件、删除、chmod、pip install 被拒绝。 -- [ ] 7.7 增加工具预算和防循环;前置条件:Tool Router core 完成;优先级:P0;验收标准:单轮最大调用数、总耗时、重复调用检测生效;测试要点:循环工具调用被截断。 -- [ ] 7.8 增加工具审计日志规划;前置条件:安全字段确认;优先级:P1;验收标准:记录工具名、风险、确认、耗时、状态、脱敏摘要;测试要点:日志不包含密钥和完整敏感输出。 +- [x] 7.1 定义 ToolCallRequest schema;前置条件:LLM tool call contract 完成;优先级:P0;验收标准:包含 name、arguments、timeout、turn id、intent;测试要点:未知字段拒绝。 +- [x] 7.2 定义 ToolDecision 和 ToolResult;前置条件:ToolCallRequest 完成;优先级:P0;验收标准:支持 execute、reject、require_confirmation、cancelled;测试要点:结果脱敏和截断。 +- [x] 7.3 实现 Tool Router fake core;前置条件:schema 完成;优先级:P0;验收标准:按工具名路由 fake adapters;测试要点:合法、非法、超时、超预算均覆盖。 +- [x] 7.4 实现工具风险分类规则;前置条件:安全策略确认;优先级:P0;验收标准:低/中/高/禁止风险可解释;测试要点:删除、上传、账号、交易、安装依赖触发高风险或禁止。 +- [x] 7.5 规划 `memory.search` 和 `memory.save`;前置条件:MemoryManager fake 完成;优先级:P0;验收标准:search 只读,save 走敏感策略;测试要点:敏感 save 需要确认或拒绝。 +- [x] 7.6 规划 `shell.readonly` adapter;前置条件:allowlist 策略确认;优先级:P0;验收标准:只允许只读命令和受限目录;测试要点:写文件、删除、chmod、pip install 被拒绝。 +- [x] 7.7 增加工具预算和防循环;前置条件:Tool Router core 完成;优先级:P0;验收标准:单轮最大调用数、总耗时、重复调用检测生效;测试要点:循环工具调用被截断。 +- [x] 7.8 增加工具审计日志规划;前置条件:安全字段确认;优先级:P1;验收标准:记录工具名、风险、确认、耗时、状态、脱敏摘要;测试要点:日志不包含密钥和完整敏感输出。 ## 8. Open Interpreter、Playwright 与电脑控制边界 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index 5f08c6d..8e6f647 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -44,6 +44,18 @@ from .full_duplex_response import ( StreamingTtsProvider, prepare_tts_sentence, ) +from .tool_router import ( + FakeToolAdapter, + MemorySaveTool, + MemorySearchTool, + ShellReadonlyTool, + ToolCallRequest, + ToolContext, + ToolDecision, + ToolResult, + ToolRouter, + ToolRiskClassifier, +) from .models import ( AudioFrame, AudioSegment, @@ -111,6 +123,16 @@ __all__ = [ "StreamingLlmProvider", "StreamingTtsProvider", "prepare_tts_sentence", + "FakeToolAdapter", + "MemorySaveTool", + "MemorySearchTool", + "ShellReadonlyTool", + "ToolCallRequest", + "ToolContext", + "ToolDecision", + "ToolResult", + "ToolRouter", + "ToolRiskClassifier", "AudioFrame", "AudioSegment", "AudioRingBuffer", diff --git a/src/owner_voice_pet/tool_router.py b/src/owner_voice_pet/tool_router.py new file mode 100644 index 0000000..c260502 --- /dev/null +++ b/src/owner_voice_pet/tool_router.py @@ -0,0 +1,314 @@ +from __future__ import annotations + +import json +import re +import shlex +import subprocess +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Literal, Protocol + +from .agent_memory import MemoryManager, MemoryRecordInput, MemoryWritePolicy + + +ToolAction = Literal["execute", "reject", "require_confirmation"] +ToolRisk = Literal["low", "medium", "high", "forbidden"] +ToolStatus = Literal["success", "failed", "cancelled", "rejected", "confirmation_required"] + + +@dataclass(frozen=True, slots=True) +class ToolCallRequest: + id: str + name: str + arguments: dict[str, object] + requested_by_turn_id: str + natural_language_intent: str = "" + timeout_ms: int = 30000 + + +@dataclass(frozen=True, slots=True) +class ToolDecision: + action: ToolAction + risk_level: ToolRisk + reason: str + sanitized_arguments: dict[str, object] = field(default_factory=dict) + confirmation_prompt: str | None = None + + +@dataclass(frozen=True, slots=True) +class ToolResult: + id: str + status: ToolStatus + output_text: str = "" + output_truncated: bool = False + error_code: str | None = None + duration_ms: int = 0 + audit_summary: str = "" + + +@dataclass(frozen=True, slots=True) +class ToolAuditRecord: + tool_name: str + risk_level: ToolRisk + action: ToolAction + status: ToolStatus + duration_ms: int + summary: str + + +@dataclass(slots=True) +class ToolContext: + memory: MemoryManager | None = None + cwd: Path = Path(".") + allowed_roots: tuple[Path, ...] = () + + +class ToolAdapter(Protocol): + name: str + + def execute(self, request: ToolCallRequest, context: ToolContext) -> ToolResult: + ... + + +class FakeToolAdapter: + def __init__(self, name: str, output: str = "ok") -> None: + self.name = name + self.output = output + self.calls: list[ToolCallRequest] = [] + + def execute(self, request: ToolCallRequest, context: ToolContext) -> ToolResult: + self.calls.append(request) + return ToolResult(request.id, "success", self.output, audit_summary=f"{self.name} executed") + + +class ToolRiskClassifier: + forbidden_patterns = ( + r"\brm\b", + r"\bchmod\b", + r"\bchown\b", + r"\bpip\s+install\b", + r"\bnpm\s+install\b", + r"\bbrew\s+install\b", + r">\s*[/\w]", + r"\bmv\b", + r"\bcp\b", + ) + high_risk_words = ( + "delete", + "upload", + "payment", + "purchase", + "trade", + "account", + "权限", + "删除", + "上传", + "支付", + "购买", + "交易", + "账号", + ) + + def classify(self, request: ToolCallRequest) -> ToolRisk: + payload = f"{request.name} {request.natural_language_intent} {json.dumps(request.arguments, ensure_ascii=False)}" + lowered = payload.lower() + if request.name == "shell.readonly": + command = str(request.arguments.get("command", "")) + if any(re.search(pattern, command, flags=re.IGNORECASE) for pattern in self.forbidden_patterns): + return "forbidden" + if any(word in lowered for word in self.high_risk_words): + return "high" + if request.name in {"memory.search", "shell.readonly"}: + return "low" + if request.name == "memory.save": + return "medium" + return "medium" + + +class ToolRouter: + def __init__( + self, + adapters: dict[str, ToolAdapter], + *, + max_calls_per_turn: int = 5, + output_limit: int = 4000, + risk_classifier: ToolRiskClassifier | None = None, + ) -> None: + self.adapters = dict(adapters) + self.max_calls_per_turn = max_calls_per_turn + self.output_limit = output_limit + self.risk_classifier = risk_classifier or ToolRiskClassifier() + self.audit_log: list[ToolAuditRecord] = [] + self._calls_by_turn: dict[str, int] = {} + self._signatures_by_turn: dict[str, set[str]] = {} + + def route(self, request: ToolCallRequest, context: ToolContext) -> ToolDecision: + if request.name not in self.adapters: + return ToolDecision("reject", "forbidden", f"unknown tool: {request.name}") + if not isinstance(request.arguments, dict): + return ToolDecision("reject", "forbidden", "tool arguments must be an object") + if self._calls_by_turn.get(request.requested_by_turn_id, 0) >= self.max_calls_per_turn: + return ToolDecision("reject", "medium", "tool call budget exceeded") + signature = self._signature(request) + if signature in self._signatures_by_turn.setdefault(request.requested_by_turn_id, set()): + return ToolDecision("reject", "medium", "duplicate tool call rejected") + risk = self.risk_classifier.classify(request) + sanitized = _sanitize_arguments(request.arguments) + if risk == "forbidden": + return ToolDecision("reject", risk, "tool request is forbidden", sanitized) + if request.name == "memory.save": + text = str(request.arguments.get("text", "")) + decision = MemoryWritePolicy().evaluate(MemoryRecordInput("fact", text)) + if decision.requires_confirmation: + return ToolDecision( + "require_confirmation", + "high", + decision.reason, + sanitized, + confirmation_prompt="是否保存这条可能敏感的长期记忆?", + ) + if risk == "high": + return ToolDecision( + "require_confirmation", + risk, + "high risk tool request requires confirmation", + sanitized, + confirmation_prompt="是否允许执行这个高风险工具请求?", + ) + return ToolDecision("execute", risk, "approved", sanitized) + + def execute(self, request: ToolCallRequest, decision: ToolDecision, context: ToolContext) -> ToolResult: + started = time.monotonic() + if decision.action == "reject": + result = ToolResult(request.id, "rejected", error_code="TOOL_REJECTED", audit_summary=decision.reason) + self._record(request, decision, result, started) + return result + if decision.action == "require_confirmation": + result = ToolResult( + request.id, + "confirmation_required", + error_code="TOOL_CONFIRMATION_REQUIRED", + audit_summary=decision.reason, + ) + self._record(request, decision, result, started) + return result + self._calls_by_turn[request.requested_by_turn_id] = self._calls_by_turn.get(request.requested_by_turn_id, 0) + 1 + self._signatures_by_turn.setdefault(request.requested_by_turn_id, set()).add(self._signature(request)) + result = self.adapters[request.name].execute(request, context) + output, truncated = _truncate(_sanitize_text(result.output_text), self.output_limit) + result = ToolResult( + id=result.id, + status=result.status, + output_text=output, + output_truncated=result.output_truncated or truncated, + error_code=result.error_code, + duration_ms=max(result.duration_ms, round((time.monotonic() - started) * 1000)), + audit_summary=_sanitize_text(result.audit_summary), + ) + self._record(request, decision, result, started) + return result + + def _record( + self, + request: ToolCallRequest, + decision: ToolDecision, + result: ToolResult, + started: float, + ) -> None: + self.audit_log.append( + ToolAuditRecord( + tool_name=request.name, + risk_level=decision.risk_level, + action=decision.action, + status=result.status, + duration_ms=max(result.duration_ms, round((time.monotonic() - started) * 1000)), + summary=_sanitize_text(result.audit_summary or decision.reason), + ) + ) + + @staticmethod + def _signature(request: ToolCallRequest) -> str: + return json.dumps( + {"name": request.name, "arguments": request.arguments}, + ensure_ascii=False, + sort_keys=True, + ) + + +class MemorySearchTool: + name = "memory.search" + + def execute(self, request: ToolCallRequest, context: ToolContext) -> ToolResult: + if context.memory is None: + return ToolResult(request.id, "failed", error_code="MEMORY_UNAVAILABLE", audit_summary="memory unavailable") + query = str(request.arguments.get("query", "")) + top_k = int(request.arguments.get("top_k", 5)) + records = context.memory.search(query, top_k=top_k) + output = "\n".join(f"- [{record.type}] {record.text}" for record in records) + return ToolResult(request.id, "success", output, audit_summary=f"returned {len(records)} memories") + + +class MemorySaveTool: + name = "memory.save" + + def execute(self, request: ToolCallRequest, context: ToolContext) -> ToolResult: + if context.memory is None: + return ToolResult(request.id, "failed", error_code="MEMORY_UNAVAILABLE", audit_summary="memory unavailable") + record_type = str(request.arguments.get("type", "fact")) + if record_type not in {"preference", "fact", "project", "task_summary"}: + return ToolResult(request.id, "failed", error_code="MEMORY_TYPE_INVALID", audit_summary="invalid memory type") + saved = context.memory.save( + MemoryRecordInput( + record_type, # type: ignore[arg-type] + str(request.arguments.get("text", "")), + metadata=dict(request.arguments.get("metadata", {}) or {}), + ) + ) + return ToolResult(request.id, "success", saved.id, audit_summary="memory saved") + + +class ShellReadonlyTool: + name = "shell.readonly" + allowed_commands = {"pwd", "ls", "find", "rg", "cat", "sed", "git"} + + def execute(self, request: ToolCallRequest, context: ToolContext) -> ToolResult: + command = str(request.arguments.get("command", "")) + parts = shlex.split(command) + if not parts or parts[0] not in self.allowed_commands: + return ToolResult(request.id, "failed", error_code="SHELL_COMMAND_NOT_ALLOWED", audit_summary="command not allowed") + if parts[0] == "git" and len(parts) > 1 and parts[1] not in {"status", "diff", "log", "show"}: + return ToolResult(request.id, "failed", error_code="SHELL_COMMAND_NOT_ALLOWED", audit_summary="git command not allowed") + completed = subprocess.run( + parts, + cwd=context.cwd, + check=False, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + timeout=max(1, int(request.timeout_ms / 1000)), + ) + status: ToolStatus = "success" if completed.returncode == 0 else "failed" + return ToolResult( + request.id, + status, + completed.stdout, + error_code=None if completed.returncode == 0 else "SHELL_COMMAND_FAILED", + audit_summary=f"readonly shell exited {completed.returncode}", + ) + + +def _sanitize_arguments(arguments: dict[str, object]) -> dict[str, object]: + return {key: _sanitize_text(str(value)) for key, value in arguments.items()} + + +def _sanitize_text(text: str) -> str: + text = re.sub(r"(?:sk|tp)-[A-Za-z0-9_\-]{16,}", "[redacted]", text) + text = re.sub(r"(?i)(authorization:\s*)\S+", r"\1[redacted]", text) + return text + + +def _truncate(text: str, limit: int) -> tuple[str, bool]: + if len(text) <= limit: + return text, False + return text[:limit] + "...[truncated]", True diff --git a/tests/test_agent_memory.py b/tests/test_agent_memory.py index 43fe332..4fab147 100644 --- a/tests/test_agent_memory.py +++ b/tests/test_agent_memory.py @@ -84,7 +84,7 @@ class AgentMemoryTests(unittest.TestCase): policy = MemoryWritePolicy() safe = policy.evaluate(MemoryRecordInput("preference", "用户喜欢 Python")) - sensitive = policy.evaluate(MemoryRecordInput("fact", "我的 api key 是 sk-abcdefghijklmnop")) + sensitive = policy.evaluate(MemoryRecordInput("fact", "我的 api key 是 secret-value")) self.assertTrue(safe.should_save) self.assertFalse(safe.requires_confirmation) diff --git a/tests/test_tool_router.py b/tests/test_tool_router.py new file mode 100644 index 0000000..bf366b8 --- /dev/null +++ b/tests/test_tool_router.py @@ -0,0 +1,135 @@ +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 sk-abcdefghijklmnop"}, + "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()