[Tool Router]:完成安全工具路由骨架,包含结构化调用、风险分类和只读工具测试
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user