Files
Owner/src/owner_voice_pet/agent_memory.py
T

372 lines
12 KiB
Python

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)