[Tool Router]:完成安全工具路由骨架,包含结构化调用、风险分类和只读工具测试

This commit is contained in:
mkbk
2026-06-18 22:08:24 +08:00
parent dc9566d8d2
commit af7eb25ba6
5 changed files with 480 additions and 9 deletions
+22
View File
@@ -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",
+314
View File
@@ -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