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