[Agent记忆工具]:完成长期记忆和工具路由闭环,包含安全确认和审计脱敏

This commit is contained in:
mkbk
2026-06-19 12:37:05 +08:00
parent 79b3e89b79
commit 81c949a7ec
8 changed files with 223 additions and 16 deletions
@@ -30,11 +30,11 @@
## 5. Memory 与 Tool Router ## 5. Memory 与 Tool Router
- [ ] 5.1 将 `AgentConversationManager` 接入 runtime;前置条件:Phase 4;优先级:P0;验收标准:memory context 注入 LLM;测试要点:相关记忆被召回。 - [x] 5.1 将 `AgentConversationManager` 接入 runtime;前置条件:Phase 4;优先级:P0;验收标准:memory context 注入 LLM;测试要点:相关记忆被召回。
- [ ] 5.2 接入 FAISS/SQLite health;前置条件:5.1;优先级:P0;验收标准:memory enabled 时检查 index;测试要点:缺失/不一致错误。 - [x] 5.2 接入 FAISS/SQLite health;前置条件:5.1;优先级:P0;验收标准:memory enabled 时检查 index;测试要点:缺失/不一致错误。
- [ ] 5.3 ToolRouter 接入 LLM tool calls;前置条件:5.1;优先级:P0;验收标准:memory.search tool result 回注入回复;测试要点:tool call integration。 - [x] 5.3 ToolRouter 接入 LLM tool calls;前置条件:5.1;优先级:P0;验收标准:memory.search tool result 回注入回复;测试要点:tool call integration。
- [ ] 5.4 高风险工具确认/拒绝;前置条件:5.3;优先级:P0;验收标准:Open Interpreter/Playwright 默认不自动执行;测试要点:高风险 fixture。 - [x] 5.4 高风险工具确认/拒绝;前置条件:5.3;优先级:P0;验收标准:Open Interpreter/Playwright 默认不自动执行;测试要点:高风险 fixture。
- [ ] 5.5 Phase 5 提交;前置条件:5.1-5.4;优先级:P0;验收标准:中文提交 `[Agent记忆工具]...`;测试要点:memory/tool/security 单测。 - [x] 5.5 Phase 5 提交;前置条件:5.1-5.4;优先级:P0;验收标准:中文提交 `[Agent记忆工具]...`;测试要点:memory/tool/security 单测。
## 6. 自我测试、文档与最终验收 ## 6. 自我测试、文档与最终验收
+1 -1
View File
@@ -296,7 +296,7 @@ class AgentConversationManager:
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class MemoryManagementPlan: class MemoryManagementPlan:
supported_commands: tuple[str, ...] = ("list", "delete", "disable", "export") supported_commands: tuple[str, ...] = ("list", "delete", "disable", "export")
default_enabled: bool = False default_enabled: bool = True
def is_sensitive_memory_text(text: str) -> bool: def is_sensitive_memory_text(text: str) -> bool:
+4 -4
View File
@@ -82,11 +82,11 @@ class AppConfig:
streaming_stt_provider: str = "faster_whisper" streaming_stt_provider: str = "faster_whisper"
streaming_stt_product_candidate: str = "sensevoice" streaming_stt_product_candidate: str = "sensevoice"
streaming_tts_provider: str = "cosyvoice" streaming_tts_provider: str = "cosyvoice"
memory_enabled: bool = False memory_enabled: bool = True
memory_provider: str = "faiss_sqlite" memory_provider: str = "faiss_sqlite"
memory_top_k: int = 5 memory_top_k: int = 5
memory_auto_save_sensitive: bool = False memory_auto_save_sensitive: bool = False
tool_router_enabled: bool = False tool_router_enabled: bool = True
tool_max_calls_per_turn: int = 5 tool_max_calls_per_turn: int = 5
tool_timeout_ms: int = 30000 tool_timeout_ms: int = 30000
openinterpreter_enabled: bool = False openinterpreter_enabled: bool = False
@@ -212,11 +212,11 @@ class AppConfig:
streaming_tts_provider=( streaming_tts_provider=(
get("STREAMING_TTS_PROVIDER", "cosyvoice") or "cosyvoice" get("STREAMING_TTS_PROVIDER", "cosyvoice") or "cosyvoice"
).lower(), ).lower(),
memory_enabled=get_bool("MEMORY_ENABLED", "0"), memory_enabled=get_bool("MEMORY_ENABLED", "1"),
memory_provider=(get("MEMORY_PROVIDER", "faiss_sqlite") or "faiss_sqlite").lower(), memory_provider=(get("MEMORY_PROVIDER", "faiss_sqlite") or "faiss_sqlite").lower(),
memory_top_k=int(get("MEMORY_TOP_K", "5") or "5"), memory_top_k=int(get("MEMORY_TOP_K", "5") or "5"),
memory_auto_save_sensitive=get_bool("MEMORY_AUTO_SAVE_SENSITIVE", "0"), memory_auto_save_sensitive=get_bool("MEMORY_AUTO_SAVE_SENSITIVE", "0"),
tool_router_enabled=get_bool("TOOL_ROUTER_ENABLED", "0"), tool_router_enabled=get_bool("TOOL_ROUTER_ENABLED", "1"),
tool_max_calls_per_turn=int(get("TOOL_MAX_CALLS_PER_TURN", "5") or "5"), tool_max_calls_per_turn=int(get("TOOL_MAX_CALLS_PER_TURN", "5") or "5"),
tool_timeout_ms=int(get("TOOL_TIMEOUT_MS", "30000") or "30000"), tool_timeout_ms=int(get("TOOL_TIMEOUT_MS", "30000") or "30000"),
openinterpreter_enabled=get_bool("OPENINTERPRETER_ENABLED", "0"), openinterpreter_enabled=get_bool("OPENINTERPRETER_ENABLED", "0"),
+100
View File
@@ -1,8 +1,18 @@
from __future__ import annotations from __future__ import annotations
import time
from dataclasses import dataclass from dataclasses import dataclass
from .agent_memory import (
AgentConversationManager,
DisabledMemoryManager,
FaissIndexManifest,
MemoryHealth,
MemoryManager,
SQLiteMemoryManager,
)
from .config import AppConfig from .config import AppConfig
from .conversation import ConversationContext
from .full_duplex_audio import AudioHub, AudioProcessingProvider, build_audio_processing_provider from .full_duplex_audio import AudioHub, AudioProcessingProvider, build_audio_processing_provider
from .full_duplex_control import CancellationGraph, FullDuplexStateMachine from .full_duplex_control import CancellationGraph, FullDuplexStateMachine
from .full_duplex_response import ( from .full_duplex_response import (
@@ -17,6 +27,15 @@ from .full_duplex_response import (
from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector
from .models import AudioFrame, Message, PipelineState from .models import AudioFrame, Message, PipelineState
from .runtime import RuntimeSummary from .runtime import RuntimeSummary
from .tool_router import (
MemorySaveTool,
MemorySearchTool,
ShellReadonlyTool,
ToolCallRequest,
ToolContext,
ToolResult,
ToolRouter,
)
@dataclass(slots=True) @dataclass(slots=True)
@@ -43,17 +62,36 @@ class FullDuplexAgentRuntime:
audio_hub: AudioHub | None = None, audio_hub: AudioHub | None = None,
llm_provider: StreamingLlmProvider | None = None, llm_provider: StreamingLlmProvider | None = None,
tts_provider: StreamingTtsProvider | None = None, tts_provider: StreamingTtsProvider | None = None,
context: ConversationContext | None = None,
memory_manager: MemoryManager | None = None,
memory_manifest: FaissIndexManifest | None = None,
tool_router: ToolRouter | None = None,
) -> None: ) -> None:
self.config = config self.config = config
self.processor = processor self.processor = processor
self.audio_hub = audio_hub self.audio_hub = audio_hub
self.llm_provider = llm_provider self.llm_provider = llm_provider
self.tts_provider = tts_provider self.tts_provider = tts_provider
self.context = context or ConversationContext(
max_messages=config.context_max_messages,
max_chars=config.context_max_chars,
)
self.memory_manager = memory_manager or self._build_memory_manager()
self.memory_manifest = memory_manifest
self.conversation_manager = AgentConversationManager(
context=self.context,
memory=self.memory_manager,
memory_enabled=config.memory_enabled,
memory_top_k=config.memory_top_k,
)
self.tool_router = tool_router or self._build_tool_router()
self.health: FullDuplexRuntimeHealth | None = None self.health: FullDuplexRuntimeHealth | None = None
self.state_machine = FullDuplexStateMachine() self.state_machine = FullDuplexStateMachine()
self.cancellation_graph = CancellationGraph("turn") self.cancellation_graph = CancellationGraph("turn")
self.interrupt_controller: InterruptController | None = None self.interrupt_controller: InterruptController | None = None
self.playback_queue = InterruptiblePlaybackQueue() self.playback_queue = InterruptiblePlaybackQueue()
self.tool_results: list[ToolResult] = []
self.tool_result_messages: list[Message] = []
def load_audio(self) -> FullDuplexRuntimeHealth: def load_audio(self) -> FullDuplexRuntimeHealth:
if self.audio_hub is None: if self.audio_hub is None:
@@ -152,6 +190,8 @@ class FullDuplexAgentRuntime:
if event.kind == "delta" and event.text_delta: if event.kind == "delta" and event.text_delta:
for sentence in segmenter.accept_delta(event.text_delta): for sentence in segmenter.accept_delta(event.text_delta):
self._synthesize_and_play_sentence(sentence, tts_session) self._synthesize_and_play_sentence(sentence, tts_session)
elif event.kind == "tool_call" and event.tool_call:
self._handle_tool_call(event.tool_call)
elif event.kind == "finish": elif event.kind == "finish":
break break
tail = segmenter.flush() tail = segmenter.flush()
@@ -162,6 +202,19 @@ class FullDuplexAgentRuntime:
self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root) self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root)
return self.playback_queue.spoken.text return self.playback_queue.spoken.text
def run_conversation_response_fixture(
self,
user_text: str,
*,
llm_events: list[LlmStreamEvent] | None = None,
) -> str:
messages = self.conversation_manager.build_messages_for_user(user_text)
spoken = self.run_streaming_response_fixture(messages, llm_events=llm_events)
self.conversation_manager.commit_user(user_text)
if spoken:
self.conversation_manager.commit_assistant(spoken)
return spoken
def _synthesize_and_play_sentence(self, sentence: str, tts_session) -> None: def _synthesize_and_play_sentence(self, sentence: str, tts_session) -> None:
if self.audio_hub is None: if self.audio_hub is None:
raise RuntimeError("audio hub is not loaded") raise RuntimeError("audio hub is not loaded")
@@ -169,6 +222,53 @@ class FullDuplexAgentRuntime:
self.playback_queue.enqueue(sentence, frames) self.playback_queue.enqueue(sentence, frames)
self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root) self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root)
def check_memory_health(self) -> MemoryHealth:
health = self.memory_manager.health_check()
errors = list(health.errors)
if self.memory_manifest is not None and hasattr(self.memory_manager, "all_records"):
records = self.memory_manager.all_records() # type: ignore[attr-defined]
errors.extend(self.memory_manifest.consistency_errors(records))
return MemoryHealth(ok=not errors, errors=tuple(errors))
def _build_memory_manager(self) -> MemoryManager:
if not self.config.memory_enabled:
return DisabledMemoryManager()
return SQLiteMemoryManager(self.config.log_dir / "memory.sqlite3")
def _build_tool_router(self) -> ToolRouter:
if not self.config.tool_router_enabled:
return ToolRouter({})
return ToolRouter(
{
"memory.search": MemorySearchTool(),
"memory.save": MemorySaveTool(),
"shell.readonly": ShellReadonlyTool(),
},
max_calls_per_turn=self.config.tool_max_calls_per_turn,
)
def _handle_tool_call(self, tool_call: dict[str, object]) -> ToolResult:
request = ToolCallRequest(
id=str(tool_call.get("id") or f"tool-{len(self.tool_results) + 1}"),
name=str(tool_call.get("name") or ""),
arguments=dict(tool_call.get("arguments") or {}),
requested_by_turn_id=str(tool_call.get("turn_id") or "turn"),
natural_language_intent=str(tool_call.get("natural_language_intent") or ""),
timeout_ms=self.config.tool_timeout_ms,
)
decision = self.tool_router.route(request, ToolContext(memory=self.memory_manager))
result = self.tool_router.execute(request, decision, ToolContext(memory=self.memory_manager))
self.tool_results.append(result)
if result.output_text:
self.tool_result_messages.append(
Message(
"system",
f"工具结果 {request.name}:\n{result.output_text}",
time.time(),
)
)
return result
def _run_audio_smoke_once(self) -> None: def _run_audio_smoke_once(self) -> None:
if self.audio_hub is None: if self.audio_hub is None:
raise RuntimeError("audio hub is not loaded") raise RuntimeError("audio hub is not loaded")
+1 -1
View File
@@ -134,7 +134,7 @@ class AgentMemoryTests(unittest.TestCase):
self.assertIn("delete", plan.supported_commands) self.assertIn("delete", plan.supported_commands)
self.assertIn("disable", plan.supported_commands) self.assertIn("disable", plan.supported_commands)
self.assertIn("export", plan.supported_commands) self.assertIn("export", plan.supported_commands)
self.assertFalse(plan.default_enabled) self.assertTrue(plan.default_enabled)
if __name__ == "__main__": if __name__ == "__main__":
+2 -2
View File
@@ -61,9 +61,9 @@ class CliAcceptanceTests(unittest.TestCase):
self.assertEqual(data["streaming_stt_provider"], "faster_whisper") self.assertEqual(data["streaming_stt_provider"], "faster_whisper")
self.assertEqual(data["streaming_stt_product_candidate"], "sensevoice") self.assertEqual(data["streaming_stt_product_candidate"], "sensevoice")
self.assertEqual(data["streaming_tts_provider"], "cosyvoice") self.assertEqual(data["streaming_tts_provider"], "cosyvoice")
self.assertFalse(data["memory_enabled"]) self.assertTrue(data["memory_enabled"])
self.assertEqual(data["memory_provider"], "faiss_sqlite") self.assertEqual(data["memory_provider"], "faiss_sqlite")
self.assertFalse(data["tool_router_enabled"]) self.assertTrue(data["tool_router_enabled"])
self.assertFalse(data["openinterpreter_enabled"]) self.assertFalse(data["openinterpreter_enabled"])
self.assertFalse(data["browser_playwright_enabled"]) self.assertFalse(data["browser_playwright_enabled"])
self.assertFalse(data["computer_control_enabled"]) self.assertFalse(data["computer_control_enabled"])
+108 -1
View File
@@ -4,7 +4,7 @@ import tempfile
import unittest import unittest
from pathlib import Path from pathlib import Path
from owner_voice_pet.agent_memory import FakeMemoryManager, MemoryRecordInput, SQLiteMemoryManager from owner_voice_pet.agent_memory import FaissIndexManifest, FakeMemoryManager, MemoryRecordInput, SQLiteMemoryManager
from owner_voice_pet.config import AppConfig from owner_voice_pet.config import AppConfig
from owner_voice_pet.full_duplex_audio import FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer from owner_voice_pet.full_duplex_audio import FakeWebRtcAudioProcessingProvider, RenderReferenceRingBuffer
from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine from owner_voice_pet.full_duplex_control import CancellationGraph, FullDuplexStateMachine
@@ -150,6 +150,113 @@ class FullDuplexIntegrationTests(unittest.TestCase):
self.assertEqual(runtime.audio_hub.render_reference.frame_count, 2) self.assertEqual(runtime.audio_hub.render_reference.frame_count, 2)
self.assertEqual(runtime.playback_queue.pending_items, 0) self.assertEqual(runtime.playback_queue.pending_items, 0)
def test_full_duplex_runtime_injects_memory_context_before_llm(self) -> None:
memory = FakeMemoryManager()
memory.save(MemoryRecordInput("preference", "用户喜欢 Python"))
llm = FakeStreamingLlmProvider([LlmStreamEvent("delta", "记住了。"), LlmStreamEvent("finish", finish_reason="stop")])
runtime = FullDuplexAgentRuntime(
config=AppConfig(audio_apm_provider="fake", audio_apm_required=False, memory_enabled=True),
memory_manager=memory,
llm_provider=llm,
tts_provider=FakeStreamingTtsProvider(),
)
spoken = runtime.run_conversation_response_fixture("Python 项目怎么做?")
self.assertEqual(spoken, "记住了。")
self.assertIn("长期记忆", llm.requests[0][1].content)
self.assertIn("用户喜欢 Python", llm.requests[0][1].content)
self.assertEqual([message.role for message in runtime.context.messages()], ["user", "assistant"])
def test_full_duplex_runtime_memory_health_detects_manifest_mismatch(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
memory = SQLiteMemoryManager(Path(tmp) / "memory.sqlite3")
saved = memory.save(MemoryRecordInput("fact", "Owner 正在做全双工语音助手"))
broken = FaissIndexManifest(
embedding_model="fake",
record_ids=(saved.id, "missing-id"),
checksums={saved.id: "wrong"},
)
runtime = FullDuplexAgentRuntime(
config=AppConfig(audio_apm_provider="fake", audio_apm_required=False, memory_enabled=True),
memory_manager=memory,
memory_manifest=broken,
)
health = runtime.check_memory_health()
self.assertFalse(health.ok)
self.assertTrue(any("missing-id" in error for error in health.errors))
self.assertTrue(any("checksum mismatch" in error for error in health.errors))
def test_full_duplex_runtime_routes_memory_search_tool_call(self) -> None:
memory = FakeMemoryManager()
memory.save(MemoryRecordInput("project", "Owner 项目正在做全双工 Agent"))
runtime = FullDuplexAgentRuntime(
config=AppConfig(
audio_apm_provider="fake",
audio_apm_required=False,
memory_enabled=True,
tool_router_enabled=True,
),
memory_manager=memory,
llm_provider=FakeStreamingLlmProvider(
[
LlmStreamEvent(
"tool_call",
tool_call={
"id": "tool-1",
"name": "memory.search",
"arguments": {"query": "Owner Agent", "top_k": 1},
"turn_id": "turn-1",
},
),
LlmStreamEvent("delta", "查到了。"),
LlmStreamEvent("finish", finish_reason="stop"),
]
),
tts_provider=FakeStreamingTtsProvider(),
)
spoken = runtime.run_conversation_response_fixture("查一下当前项目")
self.assertEqual(spoken, "查到了。")
self.assertEqual(runtime.tool_results[0].status, "success")
self.assertIn("全双工 Agent", runtime.tool_results[0].output_text)
self.assertIn("工具结果 memory.search", runtime.tool_result_messages[0].content)
def test_full_duplex_runtime_high_risk_tool_call_requires_confirmation(self) -> None:
runtime = FullDuplexAgentRuntime(
config=AppConfig(
audio_apm_provider="fake",
audio_apm_required=False,
memory_enabled=True,
tool_router_enabled=True,
),
memory_manager=FakeMemoryManager(),
llm_provider=FakeStreamingLlmProvider(
[
LlmStreamEvent(
"tool_call",
tool_call={
"id": "tool-1",
"name": "memory.search",
"arguments": {"query": "账号"},
"natural_language_intent": "上传账号资料",
"turn_id": "turn-1",
},
),
LlmStreamEvent("finish", finish_reason="stop"),
]
),
tts_provider=FakeStreamingTtsProvider(),
)
runtime.run_conversation_response_fixture("上传账号资料")
self.assertEqual(runtime.tool_results[0].status, "confirmation_required")
self.assertEqual(runtime.tool_router.audit_log[0].action, "require_confirmation")
def test_memory_restart_and_tool_search_integration(self) -> None: def test_memory_restart_and_tool_search_integration(self) -> None:
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
db_path = Path(tmp) / "memory.sqlite3" db_path = Path(tmp) / "memory.sqlite3"
+2 -2
View File
@@ -135,11 +135,11 @@ class ModelsConfigTests(unittest.TestCase):
self.assertEqual(config.streaming_stt_provider, "faster_whisper") self.assertEqual(config.streaming_stt_provider, "faster_whisper")
self.assertEqual(config.streaming_stt_product_candidate, "sensevoice") self.assertEqual(config.streaming_stt_product_candidate, "sensevoice")
self.assertEqual(config.streaming_tts_provider, "cosyvoice") self.assertEqual(config.streaming_tts_provider, "cosyvoice")
self.assertFalse(config.memory_enabled) self.assertTrue(config.memory_enabled)
self.assertEqual(config.memory_provider, "faiss_sqlite") self.assertEqual(config.memory_provider, "faiss_sqlite")
self.assertEqual(config.memory_top_k, 5) self.assertEqual(config.memory_top_k, 5)
self.assertFalse(config.memory_auto_save_sensitive) self.assertFalse(config.memory_auto_save_sensitive)
self.assertFalse(config.tool_router_enabled) self.assertTrue(config.tool_router_enabled)
self.assertEqual(config.tool_max_calls_per_turn, 5) self.assertEqual(config.tool_max_calls_per_turn, 5)
self.assertEqual(config.tool_timeout_ms, 30000) self.assertEqual(config.tool_timeout_ms, 30000)
self.assertFalse(config.openinterpreter_enabled) self.assertFalse(config.openinterpreter_enabled)