[Tool Router]:完成安全工具路由骨架,包含结构化调用、风险分类和只读工具测试
This commit is contained in:
@@ -59,14 +59,14 @@
|
|||||||
|
|
||||||
## 7. Tool Router 与安全工具
|
## 7. Tool Router 与安全工具
|
||||||
|
|
||||||
- [ ] 7.1 定义 ToolCallRequest schema;前置条件:LLM tool call contract 完成;优先级:P0;验收标准:包含 name、arguments、timeout、turn id、intent;测试要点:未知字段拒绝。
|
- [x] 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;测试要点:结果脱敏和截断。
|
- [x] 7.2 定义 ToolDecision 和 ToolResult;前置条件:ToolCallRequest 完成;优先级:P0;验收标准:支持 execute、reject、require_confirmation、cancelled;测试要点:结果脱敏和截断。
|
||||||
- [ ] 7.3 实现 Tool Router fake core;前置条件:schema 完成;优先级:P0;验收标准:按工具名路由 fake adapters;测试要点:合法、非法、超时、超预算均覆盖。
|
- [x] 7.3 实现 Tool Router fake core;前置条件:schema 完成;优先级:P0;验收标准:按工具名路由 fake adapters;测试要点:合法、非法、超时、超预算均覆盖。
|
||||||
- [ ] 7.4 实现工具风险分类规则;前置条件:安全策略确认;优先级:P0;验收标准:低/中/高/禁止风险可解释;测试要点:删除、上传、账号、交易、安装依赖触发高风险或禁止。
|
- [x] 7.4 实现工具风险分类规则;前置条件:安全策略确认;优先级:P0;验收标准:低/中/高/禁止风险可解释;测试要点:删除、上传、账号、交易、安装依赖触发高风险或禁止。
|
||||||
- [ ] 7.5 规划 `memory.search` 和 `memory.save`;前置条件:MemoryManager fake 完成;优先级:P0;验收标准:search 只读,save 走敏感策略;测试要点:敏感 save 需要确认或拒绝。
|
- [x] 7.5 规划 `memory.search` 和 `memory.save`;前置条件:MemoryManager fake 完成;优先级:P0;验收标准:search 只读,save 走敏感策略;测试要点:敏感 save 需要确认或拒绝。
|
||||||
- [ ] 7.6 规划 `shell.readonly` adapter;前置条件:allowlist 策略确认;优先级:P0;验收标准:只允许只读命令和受限目录;测试要点:写文件、删除、chmod、pip install 被拒绝。
|
- [x] 7.6 规划 `shell.readonly` adapter;前置条件:allowlist 策略确认;优先级:P0;验收标准:只允许只读命令和受限目录;测试要点:写文件、删除、chmod、pip install 被拒绝。
|
||||||
- [ ] 7.7 增加工具预算和防循环;前置条件:Tool Router core 完成;优先级:P0;验收标准:单轮最大调用数、总耗时、重复调用检测生效;测试要点:循环工具调用被截断。
|
- [x] 7.7 增加工具预算和防循环;前置条件:Tool Router core 完成;优先级:P0;验收标准:单轮最大调用数、总耗时、重复调用检测生效;测试要点:循环工具调用被截断。
|
||||||
- [ ] 7.8 增加工具审计日志规划;前置条件:安全字段确认;优先级:P1;验收标准:记录工具名、风险、确认、耗时、状态、脱敏摘要;测试要点:日志不包含密钥和完整敏感输出。
|
- [x] 7.8 增加工具审计日志规划;前置条件:安全字段确认;优先级:P1;验收标准:记录工具名、风险、确认、耗时、状态、脱敏摘要;测试要点:日志不包含密钥和完整敏感输出。
|
||||||
|
|
||||||
## 8. Open Interpreter、Playwright 与电脑控制边界
|
## 8. Open Interpreter、Playwright 与电脑控制边界
|
||||||
|
|
||||||
|
|||||||
@@ -44,6 +44,18 @@ from .full_duplex_response import (
|
|||||||
StreamingTtsProvider,
|
StreamingTtsProvider,
|
||||||
prepare_tts_sentence,
|
prepare_tts_sentence,
|
||||||
)
|
)
|
||||||
|
from .tool_router import (
|
||||||
|
FakeToolAdapter,
|
||||||
|
MemorySaveTool,
|
||||||
|
MemorySearchTool,
|
||||||
|
ShellReadonlyTool,
|
||||||
|
ToolCallRequest,
|
||||||
|
ToolContext,
|
||||||
|
ToolDecision,
|
||||||
|
ToolResult,
|
||||||
|
ToolRouter,
|
||||||
|
ToolRiskClassifier,
|
||||||
|
)
|
||||||
from .models import (
|
from .models import (
|
||||||
AudioFrame,
|
AudioFrame,
|
||||||
AudioSegment,
|
AudioSegment,
|
||||||
@@ -111,6 +123,16 @@ __all__ = [
|
|||||||
"StreamingLlmProvider",
|
"StreamingLlmProvider",
|
||||||
"StreamingTtsProvider",
|
"StreamingTtsProvider",
|
||||||
"prepare_tts_sentence",
|
"prepare_tts_sentence",
|
||||||
|
"FakeToolAdapter",
|
||||||
|
"MemorySaveTool",
|
||||||
|
"MemorySearchTool",
|
||||||
|
"ShellReadonlyTool",
|
||||||
|
"ToolCallRequest",
|
||||||
|
"ToolContext",
|
||||||
|
"ToolDecision",
|
||||||
|
"ToolResult",
|
||||||
|
"ToolRouter",
|
||||||
|
"ToolRiskClassifier",
|
||||||
"AudioFrame",
|
"AudioFrame",
|
||||||
"AudioSegment",
|
"AudioSegment",
|
||||||
"AudioRingBuffer",
|
"AudioRingBuffer",
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -84,7 +84,7 @@ class AgentMemoryTests(unittest.TestCase):
|
|||||||
policy = MemoryWritePolicy()
|
policy = MemoryWritePolicy()
|
||||||
|
|
||||||
safe = policy.evaluate(MemoryRecordInput("preference", "用户喜欢 Python"))
|
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.assertTrue(safe.should_save)
|
||||||
self.assertFalse(safe.requires_confirmation)
|
self.assertFalse(safe.requires_confirmation)
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user