[长期记忆]:完成会话记忆管理骨架,包含SQLite结构、FAISS索引校验和敏感保存策略测试
This commit is contained in:
@@ -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 与安全工具
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user