372 lines
12 KiB
Python
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)
|