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)