from __future__ import annotations import json import math import uuid from datetime import datetime, timezone from pathlib import Path from typing import Any from .config import settings from .schemas import AuditEvent, MemoryItem def utcnow() -> datetime: return datetime.now(timezone.utc) class MemoryRepository: """统一存储层:优先使用 PostgreSQL,连接失败时退回本地 JSON。""" def __init__(self) -> None: self.pool: Any | None = None self.storage_mode = "local-json" self.initialization_error: str | None = None self.path = Path(settings.data_dir) / "local-memory.json" self.memories: dict[str, list[MemoryItem]] = {} self.audit: list[AuditEvent] = [] self.embeddings: dict[tuple[str, str], dict[str, Any]] = {} async def initialize(self) -> None: """创建数据库表和索引;初始化失败不阻塞服务启动。""" try: import asyncpg # type: ignore dimension = settings.embedding_dimensions if not 1 <= dimension <= 2000: raise ValueError("EMBEDDING_DIMENSIONS 必须在 1 到 2000 之间") self.pool = await asyncpg.create_pool(settings.database_url, timeout=2) async with self.pool.acquire() as conn: await conn.execute( f""" CREATE TABLE IF NOT EXISTS memory_items ( id TEXT PRIMARY KEY, system TEXT NOT NULL, workspace_id TEXT NOT NULL, content TEXT NOT NULL, memory_type TEXT NOT NULL, scope TEXT NOT NULL, source TEXT NOT NULL, confidence DOUBLE PRECISION NOT NULL, valid_from TIMESTAMPTZ, valid_to TIMESTAMPTZ, created_at TIMESTAMPTZ NOT NULL, updated_at TIMESTAMPTZ NOT NULL, metadata JSONB NOT NULL DEFAULT '{{}}'::jsonb ); CREATE INDEX IF NOT EXISTS memory_items_lookup ON memory_items (workspace_id, system, updated_at DESC); CREATE TABLE IF NOT EXISTS audit_events ( id TEXT PRIMARY KEY, system TEXT NOT NULL, workspace_id TEXT NOT NULL, operation TEXT NOT NULL, target TEXT, status TEXT NOT NULL, details JSONB NOT NULL DEFAULT '{{}}'::jsonb, created_at TIMESTAMPTZ NOT NULL ); CREATE INDEX IF NOT EXISTS audit_events_lookup ON audit_events (workspace_id, system, created_at DESC); CREATE TABLE IF NOT EXISTS memory_embeddings ( workspace_id TEXT NOT NULL, system TEXT NOT NULL, memory_id TEXT NOT NULL, content TEXT NOT NULL, source TEXT NOT NULL, content_hash TEXT NOT NULL, embedding vector({dimension}) NOT NULL, updated_at TIMESTAMPTZ NOT NULL, PRIMARY KEY (workspace_id, system, memory_id) ); CREATE INDEX IF NOT EXISTS memory_embeddings_lookup ON memory_embeddings (workspace_id, system, updated_at DESC); CREATE INDEX IF NOT EXISTS memory_embeddings_vector_lookup ON memory_embeddings USING hnsw (embedding vector_cosine_ops); """ ) vector_type = await conn.fetchval( """ SELECT format_type(attribute.atttypid, attribute.atttypmod) FROM pg_attribute AS attribute JOIN pg_class AS relation ON relation.oid=attribute.attrelid WHERE relation.relname='memory_embeddings' AND attribute.attname='embedding' AND attribute.attnum > 0 """ ) expected_type = f"vector({dimension})" if vector_type != expected_type: raise RuntimeError( f"memory_embeddings 当前为 {vector_type},但配置要求 {expected_type}。" "请迁移或重建该向量表后重试。" ) self.storage_mode = "postgres" self.initialization_error = None except Exception as exc: if self.pool: await self.pool.close() self.pool = None self.initialization_error = str(exc) self._load_local() def _load_local(self) -> None: """从本地 JSON 恢复记忆、审计和向量数据。""" if not self.path.exists(): return try: payload = json.loads(self.path.read_text(encoding="utf-8")) self.memories = { system: [MemoryItem.model_validate(item) for item in items] for system, items in payload.get("memories", {}).items() } self.audit = [AuditEvent.model_validate(item) for item in payload.get("audit", [])] self.embeddings = { (system, memory_id): item for system, items in payload.get("embeddings", {}).items() for memory_id, item in items.items() } except Exception: self.memories = {} self.audit = [] self.embeddings = {} def _persist_local(self) -> None: self.path.parent.mkdir(parents=True, exist_ok=True) self.path.write_text( json.dumps( { "memories": { system: [item.model_dump(mode="json") for item in items] for system, items in self.memories.items() }, "audit": [event.model_dump(mode="json") for event in self.audit], "embeddings": { system: { memory_id: item for (item_system, memory_id), item in self.embeddings.items() if item_system == system } for system in {item_system for item_system, _ in self.embeddings} }, }, ensure_ascii=False, indent=2, ), encoding="utf-8", ) @staticmethod def _json_value(value: Any) -> dict[str, Any]: if isinstance(value, str): try: parsed = json.loads(value) return parsed if isinstance(parsed, dict) else {} except json.JSONDecodeError: return {} return value or {} async def add_memory(self, item: MemoryItem) -> MemoryItem: """新增或覆盖一条记忆。""" if self.pool: async with self.pool.acquire() as conn: await conn.execute( """ INSERT INTO memory_items (id, system, workspace_id, content, memory_type, scope, source, confidence, valid_from, valid_to, created_at, updated_at, metadata) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13::jsonb) ON CONFLICT (id) DO UPDATE SET system=EXCLUDED.system, workspace_id=EXCLUDED.workspace_id, content=EXCLUDED.content, memory_type=EXCLUDED.memory_type, scope=EXCLUDED.scope, source=EXCLUDED.source, confidence=EXCLUDED.confidence, valid_from=EXCLUDED.valid_from, valid_to=EXCLUDED.valid_to, updated_at=EXCLUDED.updated_at, metadata=EXCLUDED.metadata """, item.id, item.system, item.workspace_id, item.content, item.memory_type, item.scope, item.source, item.confidence, item.valid_from, item.valid_to, item.created_at, item.updated_at, json.dumps(item.metadata, ensure_ascii=False), ) else: items = self.memories.setdefault(item.system, []) items[:] = [existing for existing in items if existing.id != item.id] items.append(item) self._persist_local() return item async def add_audit(self, event: AuditEvent) -> AuditEvent: if self.pool: async with self.pool.acquire() as conn: await conn.execute( """ INSERT INTO audit_events (id, system, workspace_id, operation, target, status, details, created_at) VALUES ($1,$2,$3,$4,$5,$6,$7::jsonb,$8) """, event.id, event.system, event.workspace_id, event.operation, event.target, event.status, json.dumps(event.details, ensure_ascii=False), event.created_at, ) else: self.audit.append(event) self._persist_local() return event async def list_memories(self, system: str, query: str | None = None) -> list[MemoryItem]: if self.pool: async with self.pool.acquire() as conn: rows = await conn.fetch( """ SELECT * FROM memory_items WHERE workspace_id=$1 AND system=$2 AND ($3::text IS NULL OR content ILIKE '%' || $3 || '%') ORDER BY updated_at DESC LIMIT 100 """, settings.workspace_id, system, query or None, ) return [ MemoryItem( id=row["id"], system=row["system"], workspace_id=row["workspace_id"], content=row["content"], memory_type=row["memory_type"], scope=row["scope"], source=row["source"], confidence=row["confidence"], valid_from=row["valid_from"], valid_to=row["valid_to"], created_at=row["created_at"], updated_at=row["updated_at"], metadata=self._json_value(row["metadata"]), ) for row in rows ] items = list(self.memories.get(system, [])) if query: lowered = query.lower() items = [item for item in items if lowered in item.content.lower()] return sorted(items, key=lambda item: item.updated_at, reverse=True)[:100] async def embedding_hashes(self, system: str) -> dict[str, str]: if self.pool: async with self.pool.acquire() as conn: rows = await conn.fetch( """ SELECT memory_id, content_hash FROM memory_embeddings WHERE workspace_id=$1 AND system=$2 """, settings.workspace_id, system, ) return {row["memory_id"]: row["content_hash"] for row in rows} return { memory_id: str(item["content_hash"]) for (item_system, memory_id), item in self.embeddings.items() if item_system == system } async def upsert_embedding( self, system: str, memory_id: str, content: str, source: str, content_hash: str, embedding: list[float], ) -> None: if self.pool: async with self.pool.acquire() as conn: await conn.execute( """ INSERT INTO memory_embeddings (workspace_id, system, memory_id, content, source, content_hash, embedding, updated_at) VALUES ($1,$2,$3,$4,$5,$6,$7::vector,$8) ON CONFLICT (workspace_id, system, memory_id) DO UPDATE SET content=EXCLUDED.content, source=EXCLUDED.source, content_hash=EXCLUDED.content_hash, embedding=EXCLUDED.embedding, updated_at=EXCLUDED.updated_at """, settings.workspace_id, system, memory_id, content, source, content_hash, "[" + ",".join(str(value) for value in embedding) + "]", utcnow(), ) return self.embeddings[(system, memory_id)] = { "content": content, "source": source, "content_hash": content_hash, "embedding": embedding, "updated_at": utcnow().isoformat(), } self._persist_local() async def search_embeddings( self, system: str, query_embedding: list[float], limit: int = 5, ) -> list[dict[str, Any]]: """按余弦相似度返回向量检索结果。""" if self.pool: async with self.pool.acquire() as conn: rows = await conn.fetch( """ SELECT memory_id, content, source, 1 - (embedding <=> $3::vector) AS score FROM memory_embeddings WHERE workspace_id=$1 AND system=$2 ORDER BY embedding <=> $3::vector LIMIT $4 """, settings.workspace_id, system, "[" + ",".join(str(value) for value in query_embedding) + "]", limit, ) return [dict(row) for row in rows] def cosine(left: list[float], right: list[float]) -> float: denominator = math.sqrt(sum(value * value for value in left)) * math.sqrt( sum(value * value for value in right) ) if not denominator: return 0.0 return sum(a * b for a, b in zip(left, right)) / denominator matches = [] for (item_system, memory_id), item in self.embeddings.items(): if item_system != system: continue matches.append( { "memory_id": memory_id, "content": item["content"], "source": item["source"], "score": cosine(query_embedding, item["embedding"]), } ) return sorted(matches, key=lambda item: item["score"], reverse=True)[:limit] async def delete_embedding(self, system: str, memory_id: str) -> None: if self.pool: async with self.pool.acquire() as conn: await conn.execute( "DELETE FROM memory_embeddings WHERE workspace_id=$1 AND system=$2 AND memory_id=$3", settings.workspace_id, system, memory_id, ) return self.embeddings.pop((system, memory_id), None) self._persist_local() async def delete_memory(self, system: str, memory_id: str) -> bool: if self.pool: async with self.pool.acquire() as conn: result = await conn.execute( "DELETE FROM memory_items WHERE workspace_id=$1 AND system=$2 AND id=$3", settings.workspace_id, system, memory_id, ) return result.endswith(" 1") items = self.memories.get(system, []) remaining = [item for item in items if item.id != memory_id] deleted = len(remaining) != len(items) self.memories[system] = remaining if deleted: self._persist_local() return deleted async def list_audit(self, system: str | None = None) -> list[AuditEvent]: if self.pool: async with self.pool.acquire() as conn: rows = await conn.fetch( """ SELECT * FROM audit_events WHERE workspace_id=$1 AND ($2::text IS NULL OR system=$2) ORDER BY created_at DESC LIMIT 200 """, settings.workspace_id, system, ) return [ AuditEvent( id=row["id"], system=row["system"], workspace_id=row["workspace_id"], operation=row["operation"], target=row["target"], status=row["status"], details=self._json_value(row["details"]), created_at=row["created_at"], ) for row in rows ] events = self.audit if system is None else [event for event in self.audit if event.system == system] return sorted(events, key=lambda event: event.created_at, reverse=True)[:200] async def reset(self, system: str | None = None) -> None: """清空指定系统或当前工作区的全部数据。""" if self.pool: async with self.pool.acquire() as conn: if system: await conn.execute( "DELETE FROM memory_items WHERE workspace_id=$1 AND system=$2", settings.workspace_id, system, ) await conn.execute( "DELETE FROM memory_embeddings WHERE workspace_id=$1 AND system=$2", settings.workspace_id, system, ) await conn.execute( "DELETE FROM audit_events WHERE workspace_id=$1 AND system=$2", settings.workspace_id, system, ) else: await conn.execute("DELETE FROM memory_items WHERE workspace_id=$1", settings.workspace_id) await conn.execute("DELETE FROM memory_embeddings WHERE workspace_id=$1", settings.workspace_id) await conn.execute("DELETE FROM audit_events WHERE workspace_id=$1", settings.workspace_id) return if system: self.memories.pop(system, None) self.embeddings = { key: value for key, value in self.embeddings.items() if key[0] != system } self.audit[:] = [event for event in self.audit if event.system != system] else: self.memories.clear() self.embeddings.clear() self.audit.clear() self._persist_local() @staticmethod def memory_id(prefix: str) -> str: return f"{prefix}_{uuid.uuid4().hex[:12]}" repository = MemoryRepository()