| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466 |
- 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-first repository with a transparent local JSON fallback.
- The fallback is intentionally exposed in the API status as `local-json`.
- It is for bootstrapping the Web UI only; Docker PostgreSQL is the intended
- persistence layer for the project.
- """
- 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:
- 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()
|