[长期记忆]:完成会话记忆管理骨架,包含SQLite结构、FAISS索引校验和敏感保存策略测试

This commit is contained in:
mkbk
2026-06-18 22:02:01 +08:00
parent acdcd39e34
commit dc9566d8d2
4 changed files with 540 additions and 8 deletions
@@ -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 与安全工具
+20
View File
@@ -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",
+371
View File
@@ -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)
+141
View File
@@ -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()