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,连接失败时退回本地 JSON。"""
  14. def __init__(self) -> None:
  15. self.pool: Any | None = None
  16. self.storage_mode = "local-json"
  17. self.initialization_error: str | None = None
  18. self.path = Path(settings.data_dir) / "local-memory.json"
  19. self.memories: dict[str, list[MemoryItem]] = {}
  20. self.audit: list[AuditEvent] = []
  21. self.embeddings: dict[tuple[str, str], dict[str, Any]] = {}
  22. async def initialize(self) -> None:
  23. """创建数据库表和索引;初始化失败不阻塞服务启动。"""
  24. try:
  25. import asyncpg # type: ignore
  26. dimension = settings.embedding_dimensions
  27. if not 1 <= dimension <= 2000:
  28. raise ValueError("EMBEDDING_DIMENSIONS 必须在 1 到 2000 之间")
  29. self.pool = await asyncpg.create_pool(settings.database_url, timeout=2)
  30. async with self.pool.acquire() as conn:
  31. await conn.execute(
  32. f"""
  33. CREATE TABLE IF NOT EXISTS memory_items (
  34. id TEXT PRIMARY KEY,
  35. system TEXT NOT NULL,
  36. workspace_id TEXT NOT NULL,
  37. content TEXT NOT NULL,
  38. memory_type TEXT NOT NULL,
  39. scope TEXT NOT NULL,
  40. source TEXT NOT NULL,
  41. confidence DOUBLE PRECISION NOT NULL,
  42. valid_from TIMESTAMPTZ,
  43. valid_to TIMESTAMPTZ,
  44. created_at TIMESTAMPTZ NOT NULL,
  45. updated_at TIMESTAMPTZ NOT NULL,
  46. metadata JSONB NOT NULL DEFAULT '{{}}'::jsonb
  47. );
  48. CREATE INDEX IF NOT EXISTS memory_items_lookup
  49. ON memory_items (workspace_id, system, updated_at DESC);
  50. CREATE TABLE IF NOT EXISTS audit_events (
  51. id TEXT PRIMARY KEY,
  52. system TEXT NOT NULL,
  53. workspace_id TEXT NOT NULL,
  54. operation TEXT NOT NULL,
  55. target TEXT,
  56. status TEXT NOT NULL,
  57. details JSONB NOT NULL DEFAULT '{{}}'::jsonb,
  58. created_at TIMESTAMPTZ NOT NULL
  59. );
  60. CREATE INDEX IF NOT EXISTS audit_events_lookup
  61. ON audit_events (workspace_id, system, created_at DESC);
  62. CREATE TABLE IF NOT EXISTS memory_embeddings (
  63. workspace_id TEXT NOT NULL,
  64. system TEXT NOT NULL,
  65. memory_id TEXT NOT NULL,
  66. content TEXT NOT NULL,
  67. source TEXT NOT NULL,
  68. content_hash TEXT NOT NULL,
  69. embedding vector({dimension}) NOT NULL,
  70. updated_at TIMESTAMPTZ NOT NULL,
  71. PRIMARY KEY (workspace_id, system, memory_id)
  72. );
  73. CREATE INDEX IF NOT EXISTS memory_embeddings_lookup
  74. ON memory_embeddings (workspace_id, system, updated_at DESC);
  75. CREATE INDEX IF NOT EXISTS memory_embeddings_vector_lookup
  76. ON memory_embeddings USING hnsw (embedding vector_cosine_ops);
  77. """
  78. )
  79. vector_type = await conn.fetchval(
  80. """
  81. SELECT format_type(attribute.atttypid, attribute.atttypmod)
  82. FROM pg_attribute AS attribute
  83. JOIN pg_class AS relation ON relation.oid=attribute.attrelid
  84. WHERE relation.relname='memory_embeddings'
  85. AND attribute.attname='embedding'
  86. AND attribute.attnum > 0
  87. """
  88. )
  89. expected_type = f"vector({dimension})"
  90. if vector_type != expected_type:
  91. raise RuntimeError(
  92. f"memory_embeddings 当前为 {vector_type},但配置要求 {expected_type}。"
  93. "请迁移或重建该向量表后重试。"
  94. )
  95. self.storage_mode = "postgres"
  96. self.initialization_error = None
  97. except Exception as exc:
  98. if self.pool:
  99. await self.pool.close()
  100. self.pool = None
  101. self.initialization_error = str(exc)
  102. self._load_local()
  103. def _load_local(self) -> None:
  104. """从本地 JSON 恢复记忆、审计和向量数据。"""
  105. if not self.path.exists():
  106. return
  107. try:
  108. payload = json.loads(self.path.read_text(encoding="utf-8"))
  109. self.memories = {
  110. system: [MemoryItem.model_validate(item) for item in items]
  111. for system, items in payload.get("memories", {}).items()
  112. }
  113. self.audit = [AuditEvent.model_validate(item) for item in payload.get("audit", [])]
  114. self.embeddings = {
  115. (system, memory_id): item
  116. for system, items in payload.get("embeddings", {}).items()
  117. for memory_id, item in items.items()
  118. }
  119. except Exception:
  120. self.memories = {}
  121. self.audit = []
  122. self.embeddings = {}
  123. def _persist_local(self) -> None:
  124. self.path.parent.mkdir(parents=True, exist_ok=True)
  125. self.path.write_text(
  126. json.dumps(
  127. {
  128. "memories": {
  129. system: [item.model_dump(mode="json") for item in items]
  130. for system, items in self.memories.items()
  131. },
  132. "audit": [event.model_dump(mode="json") for event in self.audit],
  133. "embeddings": {
  134. system: {
  135. memory_id: item
  136. for (item_system, memory_id), item in self.embeddings.items()
  137. if item_system == system
  138. }
  139. for system in {item_system for item_system, _ in self.embeddings}
  140. },
  141. },
  142. ensure_ascii=False,
  143. indent=2,
  144. ),
  145. encoding="utf-8",
  146. )
  147. @staticmethod
  148. def _json_value(value: Any) -> dict[str, Any]:
  149. if isinstance(value, str):
  150. try:
  151. parsed = json.loads(value)
  152. return parsed if isinstance(parsed, dict) else {}
  153. except json.JSONDecodeError:
  154. return {}
  155. return value or {}
  156. async def add_memory(self, item: MemoryItem) -> MemoryItem:
  157. """新增或覆盖一条记忆。"""
  158. if self.pool:
  159. async with self.pool.acquire() as conn:
  160. await conn.execute(
  161. """
  162. INSERT INTO memory_items
  163. (id, system, workspace_id, content, memory_type, scope, source,
  164. confidence, valid_from, valid_to, created_at, updated_at, metadata)
  165. VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13::jsonb)
  166. ON CONFLICT (id) DO UPDATE SET
  167. system=EXCLUDED.system,
  168. workspace_id=EXCLUDED.workspace_id,
  169. content=EXCLUDED.content,
  170. memory_type=EXCLUDED.memory_type,
  171. scope=EXCLUDED.scope,
  172. source=EXCLUDED.source,
  173. confidence=EXCLUDED.confidence,
  174. valid_from=EXCLUDED.valid_from,
  175. valid_to=EXCLUDED.valid_to,
  176. updated_at=EXCLUDED.updated_at,
  177. metadata=EXCLUDED.metadata
  178. """,
  179. item.id,
  180. item.system,
  181. item.workspace_id,
  182. item.content,
  183. item.memory_type,
  184. item.scope,
  185. item.source,
  186. item.confidence,
  187. item.valid_from,
  188. item.valid_to,
  189. item.created_at,
  190. item.updated_at,
  191. json.dumps(item.metadata, ensure_ascii=False),
  192. )
  193. else:
  194. items = self.memories.setdefault(item.system, [])
  195. items[:] = [existing for existing in items if existing.id != item.id]
  196. items.append(item)
  197. self._persist_local()
  198. return item
  199. async def add_audit(self, event: AuditEvent) -> AuditEvent:
  200. if self.pool:
  201. async with self.pool.acquire() as conn:
  202. await conn.execute(
  203. """
  204. INSERT INTO audit_events
  205. (id, system, workspace_id, operation, target, status, details, created_at)
  206. VALUES ($1,$2,$3,$4,$5,$6,$7::jsonb,$8)
  207. """,
  208. event.id,
  209. event.system,
  210. event.workspace_id,
  211. event.operation,
  212. event.target,
  213. event.status,
  214. json.dumps(event.details, ensure_ascii=False),
  215. event.created_at,
  216. )
  217. else:
  218. self.audit.append(event)
  219. self._persist_local()
  220. return event
  221. async def list_memories(self, system: str, query: str | None = None) -> list[MemoryItem]:
  222. if self.pool:
  223. async with self.pool.acquire() as conn:
  224. rows = await conn.fetch(
  225. """
  226. SELECT * FROM memory_items
  227. WHERE workspace_id=$1 AND system=$2
  228. AND ($3::text IS NULL OR content ILIKE '%' || $3 || '%')
  229. ORDER BY updated_at DESC
  230. LIMIT 100
  231. """,
  232. settings.workspace_id,
  233. system,
  234. query or None,
  235. )
  236. return [
  237. MemoryItem(
  238. id=row["id"], system=row["system"], workspace_id=row["workspace_id"],
  239. content=row["content"], memory_type=row["memory_type"], scope=row["scope"],
  240. source=row["source"], confidence=row["confidence"],
  241. valid_from=row["valid_from"], valid_to=row["valid_to"],
  242. created_at=row["created_at"], updated_at=row["updated_at"],
  243. metadata=self._json_value(row["metadata"]),
  244. )
  245. for row in rows
  246. ]
  247. items = list(self.memories.get(system, []))
  248. if query:
  249. lowered = query.lower()
  250. items = [item for item in items if lowered in item.content.lower()]
  251. return sorted(items, key=lambda item: item.updated_at, reverse=True)[:100]
  252. async def embedding_hashes(self, system: str) -> dict[str, str]:
  253. if self.pool:
  254. async with self.pool.acquire() as conn:
  255. rows = await conn.fetch(
  256. """
  257. SELECT memory_id, content_hash FROM memory_embeddings
  258. WHERE workspace_id=$1 AND system=$2
  259. """,
  260. settings.workspace_id,
  261. system,
  262. )
  263. return {row["memory_id"]: row["content_hash"] for row in rows}
  264. return {
  265. memory_id: str(item["content_hash"])
  266. for (item_system, memory_id), item in self.embeddings.items()
  267. if item_system == system
  268. }
  269. async def upsert_embedding(
  270. self,
  271. system: str,
  272. memory_id: str,
  273. content: str,
  274. source: str,
  275. content_hash: str,
  276. embedding: list[float],
  277. ) -> None:
  278. if self.pool:
  279. async with self.pool.acquire() as conn:
  280. await conn.execute(
  281. """
  282. INSERT INTO memory_embeddings
  283. (workspace_id, system, memory_id, content, source, content_hash, embedding, updated_at)
  284. VALUES ($1,$2,$3,$4,$5,$6,$7::vector,$8)
  285. ON CONFLICT (workspace_id, system, memory_id) DO UPDATE SET
  286. content=EXCLUDED.content,
  287. source=EXCLUDED.source,
  288. content_hash=EXCLUDED.content_hash,
  289. embedding=EXCLUDED.embedding,
  290. updated_at=EXCLUDED.updated_at
  291. """,
  292. settings.workspace_id,
  293. system,
  294. memory_id,
  295. content,
  296. source,
  297. content_hash,
  298. "[" + ",".join(str(value) for value in embedding) + "]",
  299. utcnow(),
  300. )
  301. return
  302. self.embeddings[(system, memory_id)] = {
  303. "content": content,
  304. "source": source,
  305. "content_hash": content_hash,
  306. "embedding": embedding,
  307. "updated_at": utcnow().isoformat(),
  308. }
  309. self._persist_local()
  310. async def search_embeddings(
  311. self,
  312. system: str,
  313. query_embedding: list[float],
  314. limit: int = 5,
  315. ) -> list[dict[str, Any]]:
  316. """按余弦相似度返回向量检索结果。"""
  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. """清空指定系统或当前工作区的全部数据。"""
  405. if self.pool:
  406. async with self.pool.acquire() as conn:
  407. if system:
  408. await conn.execute(
  409. "DELETE FROM memory_items WHERE workspace_id=$1 AND system=$2",
  410. settings.workspace_id, system,
  411. )
  412. await conn.execute(
  413. "DELETE FROM memory_embeddings WHERE workspace_id=$1 AND system=$2",
  414. settings.workspace_id, system,
  415. )
  416. await conn.execute(
  417. "DELETE FROM audit_events WHERE workspace_id=$1 AND system=$2",
  418. settings.workspace_id, system,
  419. )
  420. else:
  421. await conn.execute("DELETE FROM memory_items WHERE workspace_id=$1", settings.workspace_id)
  422. await conn.execute("DELETE FROM memory_embeddings WHERE workspace_id=$1", settings.workspace_id)
  423. await conn.execute("DELETE FROM audit_events WHERE workspace_id=$1", settings.workspace_id)
  424. return
  425. if system:
  426. self.memories.pop(system, None)
  427. self.embeddings = {
  428. key: value for key, value in self.embeddings.items() if key[0] != system
  429. }
  430. self.audit[:] = [event for event in self.audit if event.system != system]
  431. else:
  432. self.memories.clear()
  433. self.embeddings.clear()
  434. self.audit.clear()
  435. self._persist_local()
  436. @staticmethod
  437. def memory_id(prefix: str) -> str:
  438. return f"{prefix}_{uuid.uuid4().hex[:12]}"
  439. repository = MemoryRepository()