db.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466
  1. from __future__ import annotations
  2. import json
  3. import math
  4. import uuid
  5. from datetime import datetime, timezone
  6. from pathlib import Path
  7. from typing import Any
  8. from .config import settings
  9. from .schemas import AuditEvent, MemoryItem
  10. def utcnow() -> datetime:
  11. return datetime.now(timezone.utc)
  12. class MemoryRepository:
  13. """PostgreSQL-first repository with a transparent local JSON fallback.
  14. The fallback is intentionally exposed in the API status as `local-json`.
  15. It is for bootstrapping the Web UI only; Docker PostgreSQL is the intended
  16. persistence layer for the project.
  17. """
  18. def __init__(self) -> None:
  19. self.pool: Any | None = None
  20. self.storage_mode = "local-json"
  21. self.initialization_error: str | None = None
  22. self.path = Path(settings.data_dir) / "local-memory.json"
  23. self.memories: dict[str, list[MemoryItem]] = {}
  24. self.audit: list[AuditEvent] = []
  25. self.embeddings: dict[tuple[str, str], dict[str, Any]] = {}
  26. async def initialize(self) -> None:
  27. try:
  28. import asyncpg # type: ignore
  29. dimension = settings.embedding_dimensions
  30. if not 1 <= dimension <= 2000:
  31. raise ValueError("EMBEDDING_DIMENSIONS 必须在 1 到 2000 之间")
  32. self.pool = await asyncpg.create_pool(settings.database_url, timeout=2)
  33. async with self.pool.acquire() as conn:
  34. await conn.execute(
  35. f"""
  36. CREATE TABLE IF NOT EXISTS memory_items (
  37. id TEXT PRIMARY KEY,
  38. system TEXT NOT NULL,
  39. workspace_id TEXT NOT NULL,
  40. content TEXT NOT NULL,
  41. memory_type TEXT NOT NULL,
  42. scope TEXT NOT NULL,
  43. source TEXT NOT NULL,
  44. confidence DOUBLE PRECISION NOT NULL,
  45. valid_from TIMESTAMPTZ,
  46. valid_to TIMESTAMPTZ,
  47. created_at TIMESTAMPTZ NOT NULL,
  48. updated_at TIMESTAMPTZ NOT NULL,
  49. metadata JSONB NOT NULL DEFAULT '{{}}'::jsonb
  50. );
  51. CREATE INDEX IF NOT EXISTS memory_items_lookup
  52. ON memory_items (workspace_id, system, updated_at DESC);
  53. CREATE TABLE IF NOT EXISTS audit_events (
  54. id TEXT PRIMARY KEY,
  55. system TEXT NOT NULL,
  56. workspace_id TEXT NOT NULL,
  57. operation TEXT NOT NULL,
  58. target TEXT,
  59. status TEXT NOT NULL,
  60. details JSONB NOT NULL DEFAULT '{{}}'::jsonb,
  61. created_at TIMESTAMPTZ NOT NULL
  62. );
  63. CREATE INDEX IF NOT EXISTS audit_events_lookup
  64. ON audit_events (workspace_id, system, created_at DESC);
  65. CREATE TABLE IF NOT EXISTS memory_embeddings (
  66. workspace_id TEXT NOT NULL,
  67. system TEXT NOT NULL,
  68. memory_id TEXT NOT NULL,
  69. content TEXT NOT NULL,
  70. source TEXT NOT NULL,
  71. content_hash TEXT NOT NULL,
  72. embedding vector({dimension}) NOT NULL,
  73. updated_at TIMESTAMPTZ NOT NULL,
  74. PRIMARY KEY (workspace_id, system, memory_id)
  75. );
  76. CREATE INDEX IF NOT EXISTS memory_embeddings_lookup
  77. ON memory_embeddings (workspace_id, system, updated_at DESC);
  78. CREATE INDEX IF NOT EXISTS memory_embeddings_vector_lookup
  79. ON memory_embeddings USING hnsw (embedding vector_cosine_ops);
  80. """
  81. )
  82. vector_type = await conn.fetchval(
  83. """
  84. SELECT format_type(attribute.atttypid, attribute.atttypmod)
  85. FROM pg_attribute AS attribute
  86. JOIN pg_class AS relation ON relation.oid=attribute.attrelid
  87. WHERE relation.relname='memory_embeddings'
  88. AND attribute.attname='embedding'
  89. AND attribute.attnum > 0
  90. """
  91. )
  92. expected_type = f"vector({dimension})"
  93. if vector_type != expected_type:
  94. raise RuntimeError(
  95. f"memory_embeddings 当前为 {vector_type},但配置要求 {expected_type}。"
  96. "请迁移或重建该向量表后重试。"
  97. )
  98. self.storage_mode = "postgres"
  99. self.initialization_error = None
  100. except Exception as exc:
  101. if self.pool:
  102. await self.pool.close()
  103. self.pool = None
  104. self.initialization_error = str(exc)
  105. self._load_local()
  106. def _load_local(self) -> None:
  107. if not self.path.exists():
  108. return
  109. try:
  110. payload = json.loads(self.path.read_text(encoding="utf-8"))
  111. self.memories = {
  112. system: [MemoryItem.model_validate(item) for item in items]
  113. for system, items in payload.get("memories", {}).items()
  114. }
  115. self.audit = [AuditEvent.model_validate(item) for item in payload.get("audit", [])]
  116. self.embeddings = {
  117. (system, memory_id): item
  118. for system, items in payload.get("embeddings", {}).items()
  119. for memory_id, item in items.items()
  120. }
  121. except Exception:
  122. self.memories = {}
  123. self.audit = []
  124. self.embeddings = {}
  125. def _persist_local(self) -> None:
  126. self.path.parent.mkdir(parents=True, exist_ok=True)
  127. self.path.write_text(
  128. json.dumps(
  129. {
  130. "memories": {
  131. system: [item.model_dump(mode="json") for item in items]
  132. for system, items in self.memories.items()
  133. },
  134. "audit": [event.model_dump(mode="json") for event in self.audit],
  135. "embeddings": {
  136. system: {
  137. memory_id: item
  138. for (item_system, memory_id), item in self.embeddings.items()
  139. if item_system == system
  140. }
  141. for system in {item_system for item_system, _ in self.embeddings}
  142. },
  143. },
  144. ensure_ascii=False,
  145. indent=2,
  146. ),
  147. encoding="utf-8",
  148. )
  149. @staticmethod
  150. def _json_value(value: Any) -> dict[str, Any]:
  151. if isinstance(value, str):
  152. try:
  153. parsed = json.loads(value)
  154. return parsed if isinstance(parsed, dict) else {}
  155. except json.JSONDecodeError:
  156. return {}
  157. return value or {}
  158. async def add_memory(self, item: MemoryItem) -> MemoryItem:
  159. if self.pool:
  160. async with self.pool.acquire() as conn:
  161. await conn.execute(
  162. """
  163. INSERT INTO memory_items
  164. (id, system, workspace_id, content, memory_type, scope, source,
  165. confidence, valid_from, valid_to, created_at, updated_at, metadata)
  166. VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13::jsonb)
  167. ON CONFLICT (id) DO UPDATE SET
  168. system=EXCLUDED.system,
  169. workspace_id=EXCLUDED.workspace_id,
  170. content=EXCLUDED.content,
  171. memory_type=EXCLUDED.memory_type,
  172. scope=EXCLUDED.scope,
  173. source=EXCLUDED.source,
  174. confidence=EXCLUDED.confidence,
  175. valid_from=EXCLUDED.valid_from,
  176. valid_to=EXCLUDED.valid_to,
  177. updated_at=EXCLUDED.updated_at,
  178. metadata=EXCLUDED.metadata
  179. """,
  180. item.id,
  181. item.system,
  182. item.workspace_id,
  183. item.content,
  184. item.memory_type,
  185. item.scope,
  186. item.source,
  187. item.confidence,
  188. item.valid_from,
  189. item.valid_to,
  190. item.created_at,
  191. item.updated_at,
  192. json.dumps(item.metadata, ensure_ascii=False),
  193. )
  194. else:
  195. items = self.memories.setdefault(item.system, [])
  196. items[:] = [existing for existing in items if existing.id != item.id]
  197. items.append(item)
  198. self._persist_local()
  199. return item
  200. async def add_audit(self, event: AuditEvent) -> AuditEvent:
  201. if self.pool:
  202. async with self.pool.acquire() as conn:
  203. await conn.execute(
  204. """
  205. INSERT INTO audit_events
  206. (id, system, workspace_id, operation, target, status, details, created_at)
  207. VALUES ($1,$2,$3,$4,$5,$6,$7::jsonb,$8)
  208. """,
  209. event.id,
  210. event.system,
  211. event.workspace_id,
  212. event.operation,
  213. event.target,
  214. event.status,
  215. json.dumps(event.details, ensure_ascii=False),
  216. event.created_at,
  217. )
  218. else:
  219. self.audit.append(event)
  220. self._persist_local()
  221. return event
  222. async def list_memories(self, system: str, query: str | None = None) -> list[MemoryItem]:
  223. if self.pool:
  224. async with self.pool.acquire() as conn:
  225. rows = await conn.fetch(
  226. """
  227. SELECT * FROM memory_items
  228. WHERE workspace_id=$1 AND system=$2
  229. AND ($3::text IS NULL OR content ILIKE '%' || $3 || '%')
  230. ORDER BY updated_at DESC
  231. LIMIT 100
  232. """,
  233. settings.workspace_id,
  234. system,
  235. query or None,
  236. )
  237. return [
  238. MemoryItem(
  239. id=row["id"], system=row["system"], workspace_id=row["workspace_id"],
  240. content=row["content"], memory_type=row["memory_type"], scope=row["scope"],
  241. source=row["source"], confidence=row["confidence"],
  242. valid_from=row["valid_from"], valid_to=row["valid_to"],
  243. created_at=row["created_at"], updated_at=row["updated_at"],
  244. metadata=self._json_value(row["metadata"]),
  245. )
  246. for row in rows
  247. ]
  248. items = list(self.memories.get(system, []))
  249. if query:
  250. lowered = query.lower()
  251. items = [item for item in items if lowered in item.content.lower()]
  252. return sorted(items, key=lambda item: item.updated_at, reverse=True)[:100]
  253. async def embedding_hashes(self, system: str) -> dict[str, str]:
  254. if self.pool:
  255. async with self.pool.acquire() as conn:
  256. rows = await conn.fetch(
  257. """
  258. SELECT memory_id, content_hash FROM memory_embeddings
  259. WHERE workspace_id=$1 AND system=$2
  260. """,
  261. settings.workspace_id,
  262. system,
  263. )
  264. return {row["memory_id"]: row["content_hash"] for row in rows}
  265. return {
  266. memory_id: str(item["content_hash"])
  267. for (item_system, memory_id), item in self.embeddings.items()
  268. if item_system == system
  269. }
  270. async def upsert_embedding(
  271. self,
  272. system: str,
  273. memory_id: str,
  274. content: str,
  275. source: str,
  276. content_hash: str,
  277. embedding: list[float],
  278. ) -> None:
  279. if self.pool:
  280. async with self.pool.acquire() as conn:
  281. await conn.execute(
  282. """
  283. INSERT INTO memory_embeddings
  284. (workspace_id, system, memory_id, content, source, content_hash, embedding, updated_at)
  285. VALUES ($1,$2,$3,$4,$5,$6,$7::vector,$8)
  286. ON CONFLICT (workspace_id, system, memory_id) DO UPDATE SET
  287. content=EXCLUDED.content,
  288. source=EXCLUDED.source,
  289. content_hash=EXCLUDED.content_hash,
  290. embedding=EXCLUDED.embedding,
  291. updated_at=EXCLUDED.updated_at
  292. """,
  293. settings.workspace_id,
  294. system,
  295. memory_id,
  296. content,
  297. source,
  298. content_hash,
  299. "[" + ",".join(str(value) for value in embedding) + "]",
  300. utcnow(),
  301. )
  302. return
  303. self.embeddings[(system, memory_id)] = {
  304. "content": content,
  305. "source": source,
  306. "content_hash": content_hash,
  307. "embedding": embedding,
  308. "updated_at": utcnow().isoformat(),
  309. }
  310. self._persist_local()
  311. async def search_embeddings(
  312. self,
  313. system: str,
  314. query_embedding: list[float],
  315. limit: int = 5,
  316. ) -> list[dict[str, Any]]:
  317. if self.pool:
  318. async with self.pool.acquire() as conn:
  319. rows = await conn.fetch(
  320. """
  321. SELECT memory_id, content, source,
  322. 1 - (embedding <=> $3::vector) AS score
  323. FROM memory_embeddings
  324. WHERE workspace_id=$1 AND system=$2
  325. ORDER BY embedding <=> $3::vector
  326. LIMIT $4
  327. """,
  328. settings.workspace_id,
  329. system,
  330. "[" + ",".join(str(value) for value in query_embedding) + "]",
  331. limit,
  332. )
  333. return [dict(row) for row in rows]
  334. def cosine(left: list[float], right: list[float]) -> float:
  335. denominator = math.sqrt(sum(value * value for value in left)) * math.sqrt(
  336. sum(value * value for value in right)
  337. )
  338. if not denominator:
  339. return 0.0
  340. return sum(a * b for a, b in zip(left, right)) / denominator
  341. matches = []
  342. for (item_system, memory_id), item in self.embeddings.items():
  343. if item_system != system:
  344. continue
  345. matches.append(
  346. {
  347. "memory_id": memory_id,
  348. "content": item["content"],
  349. "source": item["source"],
  350. "score": cosine(query_embedding, item["embedding"]),
  351. }
  352. )
  353. return sorted(matches, key=lambda item: item["score"], reverse=True)[:limit]
  354. async def delete_embedding(self, system: str, memory_id: str) -> None:
  355. if self.pool:
  356. async with self.pool.acquire() as conn:
  357. await conn.execute(
  358. "DELETE FROM memory_embeddings WHERE workspace_id=$1 AND system=$2 AND memory_id=$3",
  359. settings.workspace_id,
  360. system,
  361. memory_id,
  362. )
  363. return
  364. self.embeddings.pop((system, memory_id), None)
  365. self._persist_local()
  366. async def delete_memory(self, system: str, memory_id: str) -> bool:
  367. if self.pool:
  368. async with self.pool.acquire() as conn:
  369. result = await conn.execute(
  370. "DELETE FROM memory_items WHERE workspace_id=$1 AND system=$2 AND id=$3",
  371. settings.workspace_id, system, memory_id,
  372. )
  373. return result.endswith(" 1")
  374. items = self.memories.get(system, [])
  375. remaining = [item for item in items if item.id != memory_id]
  376. deleted = len(remaining) != len(items)
  377. self.memories[system] = remaining
  378. if deleted:
  379. self._persist_local()
  380. return deleted
  381. async def list_audit(self, system: str | None = None) -> list[AuditEvent]:
  382. if self.pool:
  383. async with self.pool.acquire() as conn:
  384. rows = await conn.fetch(
  385. """
  386. SELECT * FROM audit_events
  387. WHERE workspace_id=$1 AND ($2::text IS NULL OR system=$2)
  388. ORDER BY created_at DESC LIMIT 200
  389. """,
  390. settings.workspace_id,
  391. system,
  392. )
  393. return [
  394. AuditEvent(
  395. id=row["id"], system=row["system"], workspace_id=row["workspace_id"],
  396. operation=row["operation"], target=row["target"], status=row["status"],
  397. details=self._json_value(row["details"]), created_at=row["created_at"],
  398. )
  399. for row in rows
  400. ]
  401. events = self.audit if system is None else [event for event in self.audit if event.system == system]
  402. return sorted(events, key=lambda event: event.created_at, reverse=True)[:200]
  403. async def reset(self, system: str | None = None) -> None:
  404. if self.pool:
  405. async with self.pool.acquire() as conn:
  406. if system:
  407. await conn.execute(
  408. "DELETE FROM memory_items WHERE workspace_id=$1 AND system=$2",
  409. settings.workspace_id, system,
  410. )
  411. await conn.execute(
  412. "DELETE FROM memory_embeddings WHERE workspace_id=$1 AND system=$2",
  413. settings.workspace_id, system,
  414. )
  415. await conn.execute(
  416. "DELETE FROM audit_events WHERE workspace_id=$1 AND system=$2",
  417. settings.workspace_id, system,
  418. )
  419. else:
  420. await conn.execute("DELETE FROM memory_items WHERE workspace_id=$1", settings.workspace_id)
  421. await conn.execute("DELETE FROM memory_embeddings WHERE workspace_id=$1", settings.workspace_id)
  422. await conn.execute("DELETE FROM audit_events WHERE workspace_id=$1", settings.workspace_id)
  423. return
  424. if system:
  425. self.memories.pop(system, None)
  426. self.embeddings = {
  427. key: value for key, value in self.embeddings.items() if key[0] != system
  428. }
  429. self.audit[:] = [event for event in self.audit if event.system != system]
  430. else:
  431. self.memories.clear()
  432. self.embeddings.clear()
  433. self.audit.clear()
  434. self._persist_local()
  435. @staticmethod
  436. def memory_id(prefix: str) -> str:
  437. return f"{prefix}_{uuid.uuid4().hex[:12]}"
  438. repository = MemoryRepository()