[长期记忆]:完成会话记忆管理骨架,包含SQLite结构、FAISS索引校验和敏感保存策略测试
This commit is contained in:
@@ -48,14 +48,14 @@
|
|||||||
|
|
||||||
## 6. Conversation Manager 与长期记忆
|
## 6. Conversation Manager 与长期记忆
|
||||||
|
|
||||||
- [ ] 6.1 定义 Conversation Manager 接口;前置条件:状态机和 LLM contract 完成;优先级:P0;验收标准:协调 transcript、context、memory、LLM、tool、TTS;测试要点:正常问答流事件顺序稳定。
|
- [x] 6.1 定义 Conversation Manager 接口;前置条件:状态机和 LLM contract 完成;优先级:P0;验收标准:协调 transcript、context、memory、LLM、tool、TTS;测试要点:正常问答流事件顺序稳定。
|
||||||
- [ ] 6.2 保留短期会话上下文;前置条件:现有 ConversationContext 梳理完成;优先级:P0;验收标准:进程内历史继续支持截断;测试要点:新 runtime 不读取旧短期历史。
|
- [x] 6.2 保留短期会话上下文;前置条件:现有 ConversationContext 梳理完成;优先级:P0;验收标准:进程内历史继续支持截断;测试要点:新 runtime 不读取旧短期历史。
|
||||||
- [ ] 6.3 设计 SQLite memory schema;前置条件:记忆类型确认;优先级:P0;验收标准:包含 type、text、metadata、sensitivity、source_turn_id、checksum;测试要点:迁移创建表成功。
|
- [x] 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。
|
- [x] 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 时不读写。
|
- [x] 6.5 实现 fake MemoryManager;前置条件:接口完成;优先级:P0;验收标准:支持 search/save/delete/health_check;测试要点:关闭 memory 时不读写。
|
||||||
- [ ] 6.6 实现记忆召回链路;前置条件:fake MemoryManager 完成;优先级:P0;验收标准:final transcript 前检索 Top-K 并注入独立 memory context;测试要点:相关偏好可召回。
|
- [x] 6.6 实现记忆召回链路;前置条件:fake MemoryManager 完成;优先级:P0;验收标准:final transcript 前检索 Top-K 并注入独立 memory context;测试要点:相关偏好可召回。
|
||||||
- [ ] 6.7 实现记忆写入安全策略;前置条件:sensitivity 分类规则确认;优先级:P0;验收标准:敏感内容默认不自动保存;测试要点:API key、支付信息、账号信息被拒绝或确认。
|
- [x] 6.7 实现记忆写入安全策略;前置条件:sensitivity 分类规则确认;优先级:P0;验收标准:敏感内容默认不自动保存;测试要点:API key、支付信息、账号信息被拒绝或确认。
|
||||||
- [ ] 6.8 增加记忆管理命令规划;前置条件:schema 和安全策略完成;优先级:P2;验收标准:列出、删除、禁用、导出策略明确;测试要点:删除后检索不到。
|
- [x] 6.8 增加记忆管理命令规划;前置条件:schema 和安全策略完成;优先级:P2;验收标准:列出、删除、禁用、导出策略明确;测试要点:删除后检索不到。
|
||||||
|
|
||||||
## 7. Tool Router 与安全工具
|
## 7. Tool Router 与安全工具
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,17 @@
|
|||||||
from .config import AppConfig
|
from .config import AppConfig
|
||||||
from .assistant_pipeline import TurnController, VoiceAssistantPipeline
|
from .assistant_pipeline import TurnController, VoiceAssistantPipeline
|
||||||
from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor
|
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 .events import PipelineEvent, PipelineEventBus
|
||||||
from .full_duplex_control import (
|
from .full_duplex_control import (
|
||||||
CancellationGraph,
|
CancellationGraph,
|
||||||
@@ -66,6 +77,15 @@ __all__ = [
|
|||||||
"VoiceAssistantPipeline",
|
"VoiceAssistantPipeline",
|
||||||
"NoopAudioPreprocessor",
|
"NoopAudioPreprocessor",
|
||||||
"SherpaOnnxDenoiserPreprocessor",
|
"SherpaOnnxDenoiserPreprocessor",
|
||||||
|
"AgentConversationManager",
|
||||||
|
"DisabledMemoryManager",
|
||||||
|
"FaissIndexManifest",
|
||||||
|
"FakeMemoryManager",
|
||||||
|
"MemoryManagementPlan",
|
||||||
|
"MemoryRecord",
|
||||||
|
"MemoryRecordInput",
|
||||||
|
"MemoryWritePolicy",
|
||||||
|
"SQLiteMemoryManager",
|
||||||
"PipelineEvent",
|
"PipelineEvent",
|
||||||
"PipelineEventBus",
|
"PipelineEventBus",
|
||||||
"CancellationGraph",
|
"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