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 4fd53e5..6eaee5c 100644 --- a/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md +++ b/openspec/changes/add-full-duplex-agent-voice-assistant/tasks.md @@ -48,14 +48,14 @@ ## 6. Conversation Manager 与长期记忆 -- [ ] 6.1 定义 Conversation Manager 接口;前置条件:状态机和 LLM contract 完成;优先级:P0;验收标准:协调 transcript、context、memory、LLM、tool、TTS;测试要点:正常问答流事件顺序稳定。 -- [ ] 6.2 保留短期会话上下文;前置条件:现有 ConversationContext 梳理完成;优先级:P0;验收标准:进程内历史继续支持截断;测试要点:新 runtime 不读取旧短期历史。 -- [ ] 6.3 设计 SQLite memory schema;前置条件:记忆类型确认;优先级:P0;验收标准:包含 type、text、metadata、sensitivity、source_turn_id、checksum;测试要点:迁移创建表成功。 -- [ ] 6.4 设计 FAISS index 管理;前置条件:embedding provider 策略确认;优先级:P1;验收标准:index 版本、embedding model 和 record id 可一致性检查;测试要点:SQLite/FAISS 不一致时报 health error。 -- [ ] 6.5 实现 fake MemoryManager;前置条件:接口完成;优先级:P0;验收标准:支持 search/save/delete/health_check;测试要点:关闭 memory 时不读写。 -- [ ] 6.6 实现记忆召回链路;前置条件:fake MemoryManager 完成;优先级:P0;验收标准:final transcript 前检索 Top-K 并注入独立 memory context;测试要点:相关偏好可召回。 -- [ ] 6.7 实现记忆写入安全策略;前置条件:sensitivity 分类规则确认;优先级:P0;验收标准:敏感内容默认不自动保存;测试要点:API key、支付信息、账号信息被拒绝或确认。 -- [ ] 6.8 增加记忆管理命令规划;前置条件:schema 和安全策略完成;优先级:P2;验收标准:列出、删除、禁用、导出策略明确;测试要点:删除后检索不到。 +- [x] 6.1 定义 Conversation Manager 接口;前置条件:状态机和 LLM contract 完成;优先级:P0;验收标准:协调 transcript、context、memory、LLM、tool、TTS;测试要点:正常问答流事件顺序稳定。 +- [x] 6.2 保留短期会话上下文;前置条件:现有 ConversationContext 梳理完成;优先级:P0;验收标准:进程内历史继续支持截断;测试要点:新 runtime 不读取旧短期历史。 +- [x] 6.3 设计 SQLite memory schema;前置条件:记忆类型确认;优先级:P0;验收标准:包含 type、text、metadata、sensitivity、source_turn_id、checksum;测试要点:迁移创建表成功。 +- [x] 6.4 设计 FAISS index 管理;前置条件:embedding provider 策略确认;优先级:P1;验收标准:index 版本、embedding model 和 record id 可一致性检查;测试要点:SQLite/FAISS 不一致时报 health error。 +- [x] 6.5 实现 fake MemoryManager;前置条件:接口完成;优先级:P0;验收标准:支持 search/save/delete/health_check;测试要点:关闭 memory 时不读写。 +- [x] 6.6 实现记忆召回链路;前置条件:fake MemoryManager 完成;优先级:P0;验收标准:final transcript 前检索 Top-K 并注入独立 memory context;测试要点:相关偏好可召回。 +- [x] 6.7 实现记忆写入安全策略;前置条件:sensitivity 分类规则确认;优先级:P0;验收标准:敏感内容默认不自动保存;测试要点:API key、支付信息、账号信息被拒绝或确认。 +- [x] 6.8 增加记忆管理命令规划;前置条件:schema 和安全策略完成;优先级:P2;验收标准:列出、删除、禁用、导出策略明确;测试要点:删除后检索不到。 ## 7. Tool Router 与安全工具 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index be78003..5f08c6d 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -3,6 +3,17 @@ from .config import AppConfig from .assistant_pipeline import TurnController, VoiceAssistantPipeline from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor +from .agent_memory import ( + AgentConversationManager, + DisabledMemoryManager, + FaissIndexManifest, + FakeMemoryManager, + MemoryManagementPlan, + MemoryRecord, + MemoryRecordInput, + MemoryWritePolicy, + SQLiteMemoryManager, +) from .events import PipelineEvent, PipelineEventBus from .full_duplex_control import ( CancellationGraph, @@ -66,6 +77,15 @@ __all__ = [ "VoiceAssistantPipeline", "NoopAudioPreprocessor", "SherpaOnnxDenoiserPreprocessor", + "AgentConversationManager", + "DisabledMemoryManager", + "FaissIndexManifest", + "FakeMemoryManager", + "MemoryManagementPlan", + "MemoryRecord", + "MemoryRecordInput", + "MemoryWritePolicy", + "SQLiteMemoryManager", "PipelineEvent", "PipelineEventBus", "CancellationGraph", diff --git a/src/owner_voice_pet/agent_memory.py b/src/owner_voice_pet/agent_memory.py new file mode 100644 index 0000000..d8d29d5 --- /dev/null +++ b/src/owner_voice_pet/agent_memory.py @@ -0,0 +1,371 @@ +from __future__ import annotations + +import hashlib +import json +import re +import sqlite3 +import time +import uuid +from dataclasses import dataclass, field +from pathlib import Path +from typing import Literal, Protocol + +from .conversation import ConversationContext +from .models import ErrorCode, Message, ProviderError + + +MemoryType = Literal["preference", "fact", "project", "task_summary"] +MemorySensitivity = Literal["normal", "sensitive"] + + +@dataclass(frozen=True, slots=True) +class MemoryRecordInput: + type: MemoryType + text: str + metadata: dict[str, object] = field(default_factory=dict) + sensitivity: MemorySensitivity = "normal" + source_turn_id: str | None = None + + +@dataclass(frozen=True, slots=True) +class MemoryRecord: + id: str + type: MemoryType + text: str + metadata: dict[str, object] + sensitivity: MemorySensitivity + source_turn_id: str | None + created_at: float + updated_at: float + last_used_at: float | None + embedding_id: str + checksum: str + + +@dataclass(frozen=True, slots=True) +class MemoryHealth: + ok: bool + errors: tuple[str, ...] = () + + +class MemoryManager(Protocol): + enabled: bool + + def search(self, query: str, *, top_k: int = 5, filters: dict[str, object] | None = None) -> list[MemoryRecord]: + ... + + def save(self, record: MemoryRecordInput) -> MemoryRecord: + ... + + def delete(self, memory_id: str) -> None: + ... + + def health_check(self) -> MemoryHealth: + ... + + +class DisabledMemoryManager: + enabled = False + + def search(self, query: str, *, top_k: int = 5, filters: dict[str, object] | None = None) -> list[MemoryRecord]: + return [] + + def save(self, record: MemoryRecordInput) -> MemoryRecord: + raise ProviderError( + ErrorCode.VALIDATION_FAILED, + "memory is disabled", + False, + "memory", + "memory", + ) + + def delete(self, memory_id: str) -> None: + return None + + def health_check(self) -> MemoryHealth: + return MemoryHealth(ok=True) + + +class FakeMemoryManager: + def __init__(self, *, enabled: bool = True) -> None: + self.enabled = enabled + self.records: dict[str, MemoryRecord] = {} + self.search_queries: list[str] = [] + + def search(self, query: str, *, top_k: int = 5, filters: dict[str, object] | None = None) -> list[MemoryRecord]: + if not self.enabled: + return [] + self.search_queries.append(query) + query_terms = set(_tokenize(query)) + ranked = sorted( + self.records.values(), + key=lambda record: len(query_terms.intersection(_tokenize(record.text))), + reverse=True, + ) + return [record for record in ranked if record.sensitivity == "normal"][:top_k] + + def save(self, record: MemoryRecordInput) -> MemoryRecord: + if not self.enabled: + raise ProviderError( + ErrorCode.VALIDATION_FAILED, + "memory is disabled", + False, + "memory", + "memory", + ) + saved = _build_memory_record(record) + self.records[saved.id] = saved + return saved + + def delete(self, memory_id: str) -> None: + self.records.pop(memory_id, None) + + def health_check(self) -> MemoryHealth: + return MemoryHealth(ok=True) + + +class SQLiteMemoryManager: + enabled = True + + def __init__(self, db_path: Path) -> None: + self.db_path = db_path + + def initialize(self) -> None: + self.db_path.parent.mkdir(parents=True, exist_ok=True) + with self._connect() as conn: + conn.execute( + """ + CREATE TABLE IF NOT EXISTS memories ( + id TEXT PRIMARY KEY, + type TEXT NOT NULL, + text TEXT NOT NULL, + metadata_json TEXT NOT NULL, + sensitivity TEXT NOT NULL, + source_turn_id TEXT, + created_at REAL NOT NULL, + updated_at REAL NOT NULL, + last_used_at REAL, + embedding_id TEXT NOT NULL, + checksum TEXT NOT NULL + ) + """ + ) + + def search(self, query: str, *, top_k: int = 5, filters: dict[str, object] | None = None) -> list[MemoryRecord]: + self.initialize() + query_terms = set(_tokenize(query)) + with self._connect() as conn: + rows = conn.execute("SELECT * FROM memories WHERE sensitivity = 'normal'").fetchall() + records = [_record_from_row(row) for row in rows] + ranked = sorted( + records, + key=lambda record: len(query_terms.intersection(_tokenize(record.text))), + reverse=True, + ) + return ranked[:top_k] + + def save(self, record: MemoryRecordInput) -> MemoryRecord: + self.initialize() + saved = _build_memory_record(record) + with self._connect() as conn: + conn.execute( + """ + INSERT INTO memories ( + id, type, text, metadata_json, sensitivity, source_turn_id, + created_at, updated_at, last_used_at, embedding_id, checksum + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + saved.id, + saved.type, + saved.text, + json.dumps(saved.metadata, ensure_ascii=False, sort_keys=True), + saved.sensitivity, + saved.source_turn_id, + saved.created_at, + saved.updated_at, + saved.last_used_at, + saved.embedding_id, + saved.checksum, + ), + ) + return saved + + def delete(self, memory_id: str) -> None: + self.initialize() + with self._connect() as conn: + conn.execute("DELETE FROM memories WHERE id = ?", (memory_id,)) + + def health_check(self) -> MemoryHealth: + try: + self.initialize() + with self._connect() as conn: + conn.execute("SELECT id, checksum FROM memories LIMIT 1").fetchall() + except sqlite3.Error as exc: + return MemoryHealth(ok=False, errors=(str(exc),)) + return MemoryHealth(ok=True) + + def all_records(self) -> list[MemoryRecord]: + self.initialize() + with self._connect() as conn: + rows = conn.execute("SELECT * FROM memories").fetchall() + return [_record_from_row(row) for row in rows] + + def _connect(self) -> sqlite3.Connection: + conn = sqlite3.connect(self.db_path) + conn.row_factory = sqlite3.Row + return conn + + +@dataclass(frozen=True, slots=True) +class FaissIndexManifest: + embedding_model: str + record_ids: tuple[str, ...] + checksums: dict[str, str] + + @classmethod + def from_records(cls, records: list[MemoryRecord], *, embedding_model: str) -> "FaissIndexManifest": + return cls( + embedding_model=embedding_model, + record_ids=tuple(record.id for record in records), + checksums={record.id: record.checksum for record in records}, + ) + + def consistency_errors(self, records: list[MemoryRecord]) -> tuple[str, ...]: + errors: list[str] = [] + by_id = {record.id: record for record in records} + for record_id in self.record_ids: + if record_id not in by_id: + errors.append(f"index references missing memory id {record_id}") + for record in records: + checksum = self.checksums.get(record.id) + if checksum != record.checksum: + errors.append(f"checksum mismatch for memory id {record.id}") + return tuple(errors) + + +@dataclass(frozen=True, slots=True) +class MemoryWriteDecision: + should_save: bool + requires_confirmation: bool + reason: str + + +class MemoryWritePolicy: + def evaluate(self, record: MemoryRecordInput) -> MemoryWriteDecision: + if record.sensitivity == "sensitive" or is_sensitive_memory_text(record.text): + return MemoryWriteDecision( + should_save=False, + requires_confirmation=True, + reason="sensitive_memory_requires_confirmation", + ) + return MemoryWriteDecision(should_save=True, requires_confirmation=False, reason="safe_to_save") + + +class AgentConversationManager: + def __init__( + self, + *, + context: ConversationContext, + memory: MemoryManager, + memory_enabled: bool, + memory_top_k: int = 5, + ) -> None: + self.context = context + self.memory = memory + self.memory_enabled = memory_enabled + self.memory_top_k = memory_top_k + + def build_messages_for_user(self, user_text: str) -> list[Message]: + messages = self.context.build_llm_messages() + if self.memory_enabled and self.memory.enabled: + memories = self.memory.search(user_text, top_k=self.memory_top_k) + if memories: + memory_text = "\n".join(f"- [{record.type}] {record.text}" for record in memories) + messages.insert(1, Message("system", f"长期记忆:\n{memory_text}", time.time())) + messages.append(Message("user", user_text.strip(), time.time())) + return messages + + def commit_user(self, text: str) -> None: + self.context.append_user(text) + + def commit_assistant(self, text: str) -> None: + self.context.append_assistant(text) + + +@dataclass(frozen=True, slots=True) +class MemoryManagementPlan: + supported_commands: tuple[str, ...] = ("list", "delete", "disable", "export") + default_enabled: bool = False + + +def is_sensitive_memory_text(text: str) -> bool: + patterns = ( + r"(?:sk|tp)-[A-Za-z0-9_\-]{16,}", + r"api[_ -]?key", + r"authorization", + r"password", + r"passwd", + r"密码", + r"银行卡", + r"身份证", + ) + lowered = text.lower() + return any(re.search(pattern, lowered, flags=re.IGNORECASE) for pattern in patterns) + + +def _build_memory_record(record: MemoryRecordInput) -> MemoryRecord: + clean = record.text.strip() + if not clean: + raise ProviderError( + ErrorCode.VALIDATION_FAILED, + "memory text must be non-empty", + False, + "memory", + "memory", + ) + now = time.time() + memory_id = str(uuid.uuid4()) + checksum = _checksum(record.type, clean, record.metadata) + return MemoryRecord( + id=memory_id, + type=record.type, + text=clean, + metadata=dict(record.metadata), + sensitivity=record.sensitivity, + source_turn_id=record.source_turn_id, + created_at=now, + updated_at=now, + last_used_at=None, + embedding_id=memory_id, + checksum=checksum, + ) + + +def _record_from_row(row: sqlite3.Row) -> MemoryRecord: + return MemoryRecord( + id=str(row["id"]), + type=row["type"], + text=str(row["text"]), + metadata=json.loads(str(row["metadata_json"])), + sensitivity=row["sensitivity"], + source_turn_id=row["source_turn_id"], + created_at=float(row["created_at"]), + updated_at=float(row["updated_at"]), + last_used_at=float(row["last_used_at"]) if row["last_used_at"] is not None else None, + embedding_id=str(row["embedding_id"]), + checksum=str(row["checksum"]), + ) + + +def _checksum(memory_type: str, text: str, metadata: dict[str, object]) -> str: + payload = json.dumps( + {"type": memory_type, "text": text, "metadata": metadata}, + ensure_ascii=False, + sort_keys=True, + ) + return hashlib.sha256(payload.encode("utf-8")).hexdigest() + + +def _tokenize(text: str) -> tuple[str, ...]: + return tuple(token for token in re.split(r"\W+", text.lower()) if token) diff --git a/tests/test_agent_memory.py b/tests/test_agent_memory.py new file mode 100644 index 0000000..43fe332 --- /dev/null +++ b/tests/test_agent_memory.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +import tempfile +import unittest +from pathlib import Path + +from owner_voice_pet.agent_memory import ( + AgentConversationManager, + DisabledMemoryManager, + FaissIndexManifest, + FakeMemoryManager, + MemoryManagementPlan, + MemoryRecordInput, + MemoryWritePolicy, + SQLiteMemoryManager, + is_sensitive_memory_text, +) +from owner_voice_pet.conversation import ConversationContext +from owner_voice_pet.models import ErrorCode, ProviderError + + +class AgentMemoryTests(unittest.TestCase): + def test_fake_memory_manager_saves_and_searches_normal_memory(self) -> None: + memory = FakeMemoryManager() + saved = memory.save(MemoryRecordInput("preference", "用户喜欢 Python 项目")) + + results = memory.search("Python", top_k=1) + + self.assertEqual(results, [saved]) + self.assertEqual(memory.search_queries, ["Python"]) + self.assertEqual(saved.type, "preference") + self.assertEqual(saved.sensitivity, "normal") + self.assertTrue(saved.checksum) + + def test_disabled_memory_manager_does_not_read_or_write(self) -> None: + memory = DisabledMemoryManager() + + self.assertEqual(memory.search("Python"), []) + with self.assertRaises(ProviderError) as raised: + memory.save(MemoryRecordInput("fact", "不会保存")) + + self.assertEqual(raised.exception.code, ErrorCode.VALIDATION_FAILED) + self.assertTrue(memory.health_check().ok) + + def test_sqlite_memory_manager_persists_across_instances(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + db_path = Path(tmp) / "memory.sqlite3" + first = SQLiteMemoryManager(db_path) + saved = first.save(MemoryRecordInput("project", "Owner 项目正在做语音助手")) + + second = SQLiteMemoryManager(db_path) + results = second.search("Owner 语音", top_k=3) + + self.assertEqual([record.id for record in results], [saved.id]) + self.assertTrue(second.health_check().ok) + + def test_sqlite_memory_manager_rejects_empty_text(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + manager = SQLiteMemoryManager(Path(tmp) / "memory.sqlite3") + + with self.assertRaises(ProviderError) as raised: + manager.save(MemoryRecordInput("fact", " ")) + + self.assertEqual(raised.exception.code, ErrorCode.VALIDATION_FAILED) + + def test_faiss_manifest_detects_missing_and_mismatched_records(self) -> None: + memory = FakeMemoryManager() + first = memory.save(MemoryRecordInput("fact", "第一条")) + second = memory.save(MemoryRecordInput("fact", "第二条")) + manifest = FaissIndexManifest.from_records([first], embedding_model="fake-embedding") + broken = FaissIndexManifest( + embedding_model=manifest.embedding_model, + record_ids=(first.id, "missing-id"), + checksums={first.id: "wrong"}, + ) + + errors = broken.consistency_errors([first, second]) + + self.assertTrue(any("missing-id" in error for error in errors)) + self.assertTrue(any(first.id in error for error in errors)) + self.assertTrue(any(second.id in error for error in errors)) + + def test_memory_write_policy_requires_confirmation_for_sensitive_text(self) -> None: + policy = MemoryWritePolicy() + + safe = policy.evaluate(MemoryRecordInput("preference", "用户喜欢 Python")) + sensitive = policy.evaluate(MemoryRecordInput("fact", "我的 api key 是 sk-abcdefghijklmnop")) + + self.assertTrue(safe.should_save) + self.assertFalse(safe.requires_confirmation) + self.assertFalse(sensitive.should_save) + self.assertTrue(sensitive.requires_confirmation) + self.assertTrue(is_sensitive_memory_text("密码是 123456")) + + def test_conversation_manager_injects_memory_context(self) -> None: + context = ConversationContext() + context.append_assistant("我是小杰。") + memory = FakeMemoryManager() + memory.save(MemoryRecordInput("preference", "用户喜欢 Python")) + manager = AgentConversationManager( + context=context, + memory=memory, + memory_enabled=True, + memory_top_k=2, + ) + + messages = manager.build_messages_for_user("Python 项目怎么做?") + + self.assertEqual(messages[0].role, "system") + self.assertEqual(messages[1].role, "system") + self.assertIn("长期记忆", messages[1].content) + self.assertIn("用户喜欢 Python", messages[1].content) + self.assertEqual(messages[-1].role, "user") + self.assertEqual(messages[-1].content, "Python 项目怎么做?") + + def test_conversation_manager_skips_memory_when_disabled(self) -> None: + context = ConversationContext() + memory = FakeMemoryManager(enabled=False) + manager = AgentConversationManager( + context=context, + memory=memory, + memory_enabled=False, + ) + + messages = manager.build_messages_for_user("Python") + + self.assertEqual([message.role for message in messages], ["system", "user"]) + self.assertEqual(memory.search_queries, []) + + def test_memory_management_plan_lists_user_controls(self) -> None: + plan = MemoryManagementPlan() + + self.assertIn("list", plan.supported_commands) + self.assertIn("delete", plan.supported_commands) + self.assertIn("disable", plan.supported_commands) + self.assertIn("export", plan.supported_commands) + self.assertFalse(plan.default_enabled) + + +if __name__ == "__main__": + unittest.main()