pycontextdb 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (47) hide show
  1. contextdb/__init__.py +72 -0
  2. contextdb/agents/__init__.py +8 -0
  3. contextdb/agents/memory_bus.py +80 -0
  4. contextdb/agents/rl_manager.py +89 -0
  5. contextdb/cli.py +107 -0
  6. contextdb/client.py +516 -0
  7. contextdb/core/__init__.py +41 -0
  8. contextdb/core/config.py +89 -0
  9. contextdb/core/exceptions.py +29 -0
  10. contextdb/core/models.py +151 -0
  11. contextdb/dynamics/__init__.py +25 -0
  12. contextdb/dynamics/evolution.py +168 -0
  13. contextdb/dynamics/formation.py +193 -0
  14. contextdb/dynamics/retrieval.py +130 -0
  15. contextdb/graphs/__init__.py +17 -0
  16. contextdb/graphs/base.py +46 -0
  17. contextdb/graphs/causal.py +224 -0
  18. contextdb/graphs/entity.py +251 -0
  19. contextdb/graphs/semantic.py +156 -0
  20. contextdb/graphs/temporal.py +173 -0
  21. contextdb/integrations/__init__.py +10 -0
  22. contextdb/integrations/autogen.py +39 -0
  23. contextdb/integrations/crewai.py +41 -0
  24. contextdb/integrations/langchain.py +132 -0
  25. contextdb/integrations/openai_tools.py +124 -0
  26. contextdb/memory/__init__.py +9 -0
  27. contextdb/memory/experiential.py +102 -0
  28. contextdb/memory/factual.py +58 -0
  29. contextdb/memory/working.py +90 -0
  30. contextdb/privacy/__init__.py +9 -0
  31. contextdb/privacy/audit.py +199 -0
  32. contextdb/privacy/pii_detector.py +173 -0
  33. contextdb/privacy/retention.py +99 -0
  34. contextdb/py.typed +0 -0
  35. contextdb/store/__init__.py +15 -0
  36. contextdb/store/base.py +67 -0
  37. contextdb/store/sqlite_store.py +517 -0
  38. contextdb/store/vector_index.py +241 -0
  39. contextdb/utils/__init__.py +22 -0
  40. contextdb/utils/embeddings.py +159 -0
  41. contextdb/utils/llm.py +139 -0
  42. contextdb/utils/migrations.py +159 -0
  43. pycontextdb-0.1.0.dist-info/METADATA +589 -0
  44. pycontextdb-0.1.0.dist-info/RECORD +47 -0
  45. pycontextdb-0.1.0.dist-info/WHEEL +4 -0
  46. pycontextdb-0.1.0.dist-info/entry_points.txt +2 -0
  47. pycontextdb-0.1.0.dist-info/licenses/LICENSE +190 -0
@@ -0,0 +1,46 @@
1
+ """Shared scaffolding for graph implementations.
2
+
3
+ Each graph stores its edges in its own SQLite table (same file as the memory
4
+ store) so the free-tier deployment stays zero-config. :class:`BaseGraph`
5
+ defines the contract and supplies a small helper for table creation.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from abc import ABC, abstractmethod
11
+ from typing import TYPE_CHECKING, Any
12
+
13
+ if TYPE_CHECKING:
14
+ from contextdb.core.models import Edge
15
+ from contextdb.store.sqlite_store import SQLiteStore
16
+
17
+
18
+ class BaseGraph(ABC):
19
+ """Common contract for all graphs over memory items."""
20
+
21
+ def __init__(self, store: SQLiteStore) -> None:
22
+ self.store = store
23
+
24
+ @abstractmethod
25
+ async def initialize(self) -> None:
26
+ """Create tables / indices. Idempotent."""
27
+
28
+ @abstractmethod
29
+ async def add_node(self, memory_id: str, data: dict[str, Any]) -> None: ...
30
+
31
+ @abstractmethod
32
+ async def add_edge(self, edge: Edge) -> None: ...
33
+
34
+ @abstractmethod
35
+ async def get_neighbors(
36
+ self,
37
+ memory_id: str,
38
+ depth: int = 1,
39
+ max_results: int = 20,
40
+ ) -> list[tuple[str, float]]: ...
41
+
42
+ @abstractmethod
43
+ async def remove_node(self, memory_id: str) -> None: ...
44
+
45
+ @abstractmethod
46
+ async def get_edges(self, memory_id: str) -> list[Edge]: ...
@@ -0,0 +1,224 @@
1
+ """Causal graph — LLM-inferred cause/effect links between memories.
2
+
3
+ The causal relationship is inherently a judgement call, so we ask the LLM
4
+ for a structured verdict over a pair of memories. Edges carry a confidence
5
+ weight in ``[0, 1]`` and a free-form ``reasoning`` payload.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+ from datetime import datetime
12
+ from typing import TYPE_CHECKING, Any
13
+
14
+ from contextdb.core.models import Edge
15
+ from contextdb.graphs.base import BaseGraph
16
+
17
+ if TYPE_CHECKING:
18
+ from contextdb.store.sqlite_store import SQLiteStore
19
+ from contextdb.utils.llm import LLMProvider
20
+
21
+ _SCHEMA = """
22
+ CREATE TABLE IF NOT EXISTS causal_edges (
23
+ source_id TEXT NOT NULL,
24
+ target_id TEXT NOT NULL,
25
+ relation TEXT NOT NULL,
26
+ weight REAL NOT NULL,
27
+ reasoning TEXT DEFAULT '',
28
+ metadata TEXT DEFAULT '{}',
29
+ created_at TEXT NOT NULL,
30
+ PRIMARY KEY (source_id, target_id)
31
+ );
32
+ CREATE INDEX IF NOT EXISTS idx_causal_source ON causal_edges(source_id);
33
+ CREATE INDEX IF NOT EXISTS idx_causal_target ON causal_edges(target_id);
34
+ """
35
+
36
+ _INFER_PROMPT = """You are a reasoning engine. Decide whether memory A caused,
37
+ enabled, or has no causal link to memory B. Return strict JSON.
38
+
39
+ Schema:
40
+ {"relation": "CAUSES|ENABLES|NONE", "confidence": 0.0-1.0, "reasoning": "string"}
41
+
42
+ Memory A (earlier): "{a_content}"
43
+ Memory B (later): "{b_content}"
44
+ """
45
+
46
+
47
+ def _safe_json(text: str) -> dict[str, Any]:
48
+ text = text.strip()
49
+ if text.startswith("```"):
50
+ lines = text.splitlines()
51
+ text = "\n".join(line for line in lines if not line.startswith("```"))
52
+ try:
53
+ loaded = json.loads(text)
54
+ return loaded if isinstance(loaded, dict) else {}
55
+ except json.JSONDecodeError:
56
+ start = text.find("{")
57
+ end = text.rfind("}")
58
+ if start != -1 and end != -1 and end > start:
59
+ try:
60
+ loaded = json.loads(text[start : end + 1])
61
+ return loaded if isinstance(loaded, dict) else {}
62
+ except json.JSONDecodeError:
63
+ return {}
64
+ return {}
65
+
66
+
67
+ class CausalGraph(BaseGraph):
68
+ """LLM-inferred cause/effect edges over memory items."""
69
+
70
+ def __init__(
71
+ self,
72
+ store: SQLiteStore,
73
+ llm: LLMProvider,
74
+ confidence_threshold: float = 0.5,
75
+ candidate_window: int = 10,
76
+ ) -> None:
77
+ super().__init__(store)
78
+ self.llm = llm
79
+ self.confidence_threshold = confidence_threshold
80
+ self.candidate_window = candidate_window
81
+
82
+ async def initialize(self) -> None:
83
+ conn = self.store._require_conn()
84
+ await conn.executescript(_SCHEMA)
85
+ await conn.commit()
86
+
87
+ async def add_node(self, memory_id: str, data: dict[str, Any]) -> None:
88
+ content = data.get("content", "")
89
+ if not content:
90
+ return
91
+ candidates = await self._recent_candidates(memory_id)
92
+ for other_id, other_content in candidates:
93
+ verdict = await self._infer(other_content, content)
94
+ relation = str(verdict.get("relation", "NONE")).upper()
95
+ confidence = float(verdict.get("confidence", 0.0))
96
+ if relation == "NONE" or confidence < self.confidence_threshold:
97
+ continue
98
+ reasoning = str(verdict.get("reasoning", ""))
99
+ await self.add_edge(
100
+ Edge(
101
+ source_id=other_id,
102
+ target_id=memory_id,
103
+ graph_type="causal",
104
+ weight=confidence,
105
+ metadata={"relation": relation, "reasoning": reasoning},
106
+ )
107
+ )
108
+
109
+ async def _infer(self, a_content: str, b_content: str) -> dict[str, Any]:
110
+ prompt = _INFER_PROMPT.replace("{a_content}", a_content).replace(
111
+ "{b_content}", b_content
112
+ )
113
+ response = await self.llm.generate(prompt, temperature=0.0, max_tokens=200)
114
+ return _safe_json(response)
115
+
116
+ async def _recent_candidates(self, memory_id: str) -> list[tuple[str, str]]:
117
+ conn = self.store._require_conn()
118
+ cursor = await conn.execute(
119
+ "SELECT id, content FROM memories WHERE id != ? AND status = 'ACTIVE' "
120
+ "ORDER BY created_at DESC LIMIT ?",
121
+ (memory_id, self.candidate_window),
122
+ )
123
+ rows = await cursor.fetchall()
124
+ return [(row["id"], row["content"]) for row in rows]
125
+
126
+ async def add_edge(self, edge: Edge) -> None:
127
+ conn = self.store._require_conn()
128
+ meta = edge.metadata or {}
129
+ await conn.execute(
130
+ "INSERT OR REPLACE INTO causal_edges "
131
+ "(source_id, target_id, relation, weight, reasoning, metadata, created_at) "
132
+ "VALUES (?,?,?,?,?,?,?)",
133
+ (
134
+ edge.source_id,
135
+ edge.target_id,
136
+ str(meta.get("relation", "CAUSES")),
137
+ edge.weight,
138
+ str(meta.get("reasoning", "")),
139
+ json.dumps(meta),
140
+ edge.created_at.isoformat(),
141
+ ),
142
+ )
143
+ await conn.commit()
144
+
145
+ async def remove_node(self, memory_id: str) -> None:
146
+ conn = self.store._require_conn()
147
+ await conn.execute(
148
+ "DELETE FROM causal_edges WHERE source_id = ? OR target_id = ?",
149
+ (memory_id, memory_id),
150
+ )
151
+ await conn.commit()
152
+
153
+ async def get_edges(self, memory_id: str) -> list[Edge]:
154
+ conn = self.store._require_conn()
155
+ cursor = await conn.execute(
156
+ "SELECT source_id, target_id, relation, weight, reasoning, metadata, created_at "
157
+ "FROM causal_edges WHERE source_id = ? OR target_id = ?",
158
+ (memory_id, memory_id),
159
+ )
160
+ rows = await cursor.fetchall()
161
+ out: list[Edge] = []
162
+ for row in rows:
163
+ meta = json.loads(row["metadata"] or "{}")
164
+ meta.setdefault("relation", row["relation"])
165
+ meta.setdefault("reasoning", row["reasoning"])
166
+ out.append(
167
+ Edge(
168
+ source_id=row["source_id"],
169
+ target_id=row["target_id"],
170
+ graph_type="causal",
171
+ weight=float(row["weight"]),
172
+ metadata=meta,
173
+ created_at=datetime.fromisoformat(row["created_at"]),
174
+ )
175
+ )
176
+ return out
177
+
178
+ async def get_neighbors(
179
+ self,
180
+ memory_id: str,
181
+ depth: int = 1,
182
+ max_results: int = 20,
183
+ ) -> list[tuple[str, float]]:
184
+ conn = self.store._require_conn()
185
+ visited: dict[str, float] = {memory_id: 0.0}
186
+ frontier: list[tuple[str, float]] = [(memory_id, 1.0)]
187
+ for _ in range(depth):
188
+ next_frontier: list[tuple[str, float]] = []
189
+ for node, cum in frontier:
190
+ cursor = await conn.execute(
191
+ "SELECT target_id, weight FROM causal_edges WHERE source_id = ?",
192
+ (node,),
193
+ )
194
+ rows = await cursor.fetchall()
195
+ for row in rows:
196
+ tid = row["target_id"]
197
+ w = cum * float(row["weight"])
198
+ if tid in visited and visited[tid] >= w:
199
+ continue
200
+ visited[tid] = w
201
+ next_frontier.append((tid, w))
202
+ frontier = next_frontier
203
+ visited.pop(memory_id, None)
204
+ return sorted(visited.items(), key=lambda kv: kv[1], reverse=True)[:max_results]
205
+
206
+ async def get_causal_chain(self, memory_id: str, max_depth: int = 5) -> list[str]:
207
+ """Return the longest cause→effect chain ending at ``memory_id``."""
208
+ conn = self.store._require_conn()
209
+ chain = [memory_id]
210
+ seen = {memory_id}
211
+ current = memory_id
212
+ for _ in range(max_depth):
213
+ cursor = await conn.execute(
214
+ "SELECT source_id FROM causal_edges "
215
+ "WHERE target_id = ? ORDER BY weight DESC LIMIT 1",
216
+ (current,),
217
+ )
218
+ row = await cursor.fetchone()
219
+ if row is None or row["source_id"] in seen:
220
+ break
221
+ current = row["source_id"]
222
+ chain.append(current)
223
+ seen.add(current)
224
+ return list(reversed(chain))
@@ -0,0 +1,251 @@
1
+ """Entity graph — memories and named entities linked bidirectionally.
2
+
3
+ Entity extraction is delegated to the configured LLM. Responses are expected
4
+ to be JSON; we tolerate minor formatting noise and deduplicate on lowercased
5
+ entity names.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+ from datetime import datetime, timezone
12
+ from typing import TYPE_CHECKING, Any
13
+ from uuid import uuid4
14
+
15
+ from contextdb.core.models import Edge, Entity
16
+ from contextdb.graphs.base import BaseGraph
17
+
18
+ if TYPE_CHECKING:
19
+ from contextdb.store.sqlite_store import SQLiteStore
20
+ from contextdb.utils.llm import LLMProvider
21
+
22
+ _SCHEMA = """
23
+ CREATE TABLE IF NOT EXISTS entities (
24
+ id TEXT PRIMARY KEY,
25
+ name TEXT NOT NULL,
26
+ entity_type TEXT NOT NULL,
27
+ attributes TEXT DEFAULT '{}',
28
+ created_at TEXT NOT NULL,
29
+ updated_at TEXT NOT NULL
30
+ );
31
+ CREATE UNIQUE INDEX IF NOT EXISTS idx_entities_name_lc ON entities(LOWER(name));
32
+
33
+ CREATE TABLE IF NOT EXISTS memory_entity_edges (
34
+ memory_id TEXT NOT NULL,
35
+ entity_id TEXT NOT NULL,
36
+ relation TEXT DEFAULT 'MENTIONS',
37
+ created_at TEXT NOT NULL,
38
+ PRIMARY KEY (memory_id, entity_id)
39
+ );
40
+ CREATE INDEX IF NOT EXISTS idx_mee_memory ON memory_entity_edges(memory_id);
41
+ CREATE INDEX IF NOT EXISTS idx_mee_entity ON memory_entity_edges(entity_id);
42
+ """
43
+
44
+ _EXTRACT_PROMPT = """Extract all named entities from the text. Return strict JSON.
45
+
46
+ Schema:
47
+ {"entities": [
48
+ {"name": "string",
49
+ "type": "PERSON|ORG|PRODUCT|LOCATION|EVENT|OTHER",
50
+ "attributes": {}}
51
+ ]}
52
+
53
+ Text:
54
+ "{text}"
55
+ """
56
+
57
+
58
+ def _safe_json(text: str) -> dict[str, Any]:
59
+ text = text.strip()
60
+ if text.startswith("```"):
61
+ # strip triple-backtick fences
62
+ lines = text.splitlines()
63
+ text = "\n".join(line for line in lines if not line.startswith("```"))
64
+ try:
65
+ loaded = json.loads(text)
66
+ return loaded if isinstance(loaded, dict) else {}
67
+ except json.JSONDecodeError:
68
+ start = text.find("{")
69
+ end = text.rfind("}")
70
+ if start != -1 and end != -1 and end > start:
71
+ try:
72
+ loaded = json.loads(text[start : end + 1])
73
+ return loaded if isinstance(loaded, dict) else {}
74
+ except json.JSONDecodeError:
75
+ return {}
76
+ return {}
77
+
78
+
79
+ class EntityGraph(BaseGraph):
80
+ """LLM-extracted entity overlay over the memory store."""
81
+
82
+ def __init__(self, store: SQLiteStore, llm: LLMProvider) -> None:
83
+ super().__init__(store)
84
+ self.llm = llm
85
+
86
+ async def initialize(self) -> None:
87
+ conn = self.store._require_conn()
88
+ await conn.executescript(_SCHEMA)
89
+ await conn.commit()
90
+
91
+ async def extract_entities(self, content: str) -> list[Entity]:
92
+ response = await self.llm.generate(_EXTRACT_PROMPT.replace("{text}", content))
93
+ payload = _safe_json(response)
94
+ out: list[Entity] = []
95
+ for raw in payload.get("entities", []) or []:
96
+ name = str(raw.get("name", "")).strip()
97
+ if not name:
98
+ continue
99
+ out.append(
100
+ Entity(
101
+ name=name,
102
+ entity_type=str(raw.get("type", "OTHER")),
103
+ attributes=raw.get("attributes") or {},
104
+ )
105
+ )
106
+ return out
107
+
108
+ async def _find_or_create(self, entity: Entity) -> str:
109
+ conn = self.store._require_conn()
110
+ cursor = await conn.execute(
111
+ "SELECT id, attributes FROM entities WHERE LOWER(name) = ?",
112
+ (entity.name.lower(),),
113
+ )
114
+ row = await cursor.fetchone()
115
+ if row is not None:
116
+ if entity.attributes:
117
+ existing = json.loads(row["attributes"] or "{}")
118
+ existing.update(entity.attributes)
119
+ await conn.execute(
120
+ "UPDATE entities SET attributes = ?, updated_at = ? WHERE id = ?",
121
+ (
122
+ json.dumps(existing),
123
+ datetime.now(tz=timezone.utc).isoformat(),
124
+ row["id"],
125
+ ),
126
+ )
127
+ await conn.commit()
128
+ return str(row["id"])
129
+ eid = str(uuid4())
130
+ now = datetime.now(tz=timezone.utc).isoformat()
131
+ await conn.execute(
132
+ "INSERT INTO entities (id, name, entity_type, attributes, created_at, updated_at) "
133
+ "VALUES (?,?,?,?,?,?)",
134
+ (eid, entity.name, entity.entity_type, json.dumps(entity.attributes), now, now),
135
+ )
136
+ await conn.commit()
137
+ return eid
138
+
139
+ async def _link_memory(self, memory_id: str, entity_id: str) -> None:
140
+ conn = self.store._require_conn()
141
+ await conn.execute(
142
+ "INSERT OR IGNORE INTO memory_entity_edges "
143
+ "(memory_id, entity_id, relation, created_at) VALUES (?,?,?,?)",
144
+ (memory_id, entity_id, "MENTIONS", datetime.now(tz=timezone.utc).isoformat()),
145
+ )
146
+ await conn.commit()
147
+
148
+ async def add_node(self, memory_id: str, data: dict[str, Any]) -> None:
149
+ content = data.get("content", "")
150
+ if not content:
151
+ return
152
+ extracted = await self.extract_entities(content)
153
+ if not extracted:
154
+ return
155
+ for ent in extracted:
156
+ eid = await self._find_or_create(ent)
157
+ await self._link_memory(memory_id, eid)
158
+
159
+ async def add_edge(self, edge: Edge) -> None:
160
+ # Entity graph's primary edges are memory↔entity; generic Edge is a no-op.
161
+ return None
162
+
163
+ async def remove_node(self, memory_id: str) -> None:
164
+ conn = self.store._require_conn()
165
+ await conn.execute(
166
+ "DELETE FROM memory_entity_edges WHERE memory_id = ?", (memory_id,)
167
+ )
168
+ await conn.commit()
169
+
170
+ async def get_edges(self, memory_id: str) -> list[Edge]:
171
+ conn = self.store._require_conn()
172
+ cursor = await conn.execute(
173
+ "SELECT entity_id, created_at FROM memory_entity_edges WHERE memory_id = ?",
174
+ (memory_id,),
175
+ )
176
+ rows = await cursor.fetchall()
177
+ return [
178
+ Edge(
179
+ source_id=memory_id,
180
+ target_id=row["entity_id"],
181
+ graph_type="entity",
182
+ weight=1.0,
183
+ created_at=datetime.fromisoformat(row["created_at"]),
184
+ )
185
+ for row in rows
186
+ ]
187
+
188
+ async def get_neighbors(
189
+ self,
190
+ memory_id: str,
191
+ depth: int = 1,
192
+ max_results: int = 20,
193
+ ) -> list[tuple[str, float]]:
194
+ conn = self.store._require_conn()
195
+ cursor = await conn.execute(
196
+ "SELECT entity_id FROM memory_entity_edges WHERE memory_id = ?", (memory_id,)
197
+ )
198
+ entity_ids = [row["entity_id"] for row in await cursor.fetchall()]
199
+ if not entity_ids:
200
+ return []
201
+ placeholders = ",".join(["?"] * len(entity_ids))
202
+ cursor = await conn.execute(
203
+ f"SELECT memory_id, COUNT(*) as shared "
204
+ f"FROM memory_entity_edges WHERE entity_id IN ({placeholders}) AND memory_id != ? "
205
+ "GROUP BY memory_id ORDER BY shared DESC LIMIT ?",
206
+ [*entity_ids, memory_id, max_results],
207
+ )
208
+ rows = await cursor.fetchall()
209
+ return [(row["memory_id"], float(row["shared"])) for row in rows]
210
+
211
+ async def get_entity_profile(self, name: str) -> dict[str, Any]:
212
+ conn = self.store._require_conn()
213
+ cursor = await conn.execute(
214
+ "SELECT id, name, entity_type, attributes FROM entities WHERE LOWER(name) = ?",
215
+ (name.lower(),),
216
+ )
217
+ row = await cursor.fetchone()
218
+ if row is None:
219
+ return {"name": name, "memories": [], "attributes": {}}
220
+ eid = row["id"]
221
+ mem_cursor = await conn.execute(
222
+ "SELECT memory_id FROM memory_entity_edges WHERE entity_id = ?", (eid,)
223
+ )
224
+ memory_ids = [r["memory_id"] for r in await mem_cursor.fetchall()]
225
+ return {
226
+ "id": eid,
227
+ "name": row["name"],
228
+ "entity_type": row["entity_type"],
229
+ "attributes": json.loads(row["attributes"] or "{}"),
230
+ "memories": memory_ids,
231
+ }
232
+
233
+ async def get_entity_memories(self, name: str) -> list[str]:
234
+ profile = await self.get_entity_profile(name)
235
+ return list(profile.get("memories", []))
236
+
237
+ async def list_entities(self) -> list[Entity]:
238
+ conn = self.store._require_conn()
239
+ cursor = await conn.execute(
240
+ "SELECT id, name, entity_type, attributes FROM entities ORDER BY name"
241
+ )
242
+ rows = await cursor.fetchall()
243
+ return [
244
+ Entity(
245
+ name=row["name"],
246
+ entity_type=row["entity_type"],
247
+ attributes=json.loads(row["attributes"] or "{}"),
248
+ memory_ids=await self.get_entity_memories(row["name"]),
249
+ )
250
+ for row in rows
251
+ ]
@@ -0,0 +1,156 @@
1
+ """Semantic graph — edges between memories with high embedding similarity."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ from datetime import datetime
7
+ from typing import TYPE_CHECKING, Any
8
+
9
+ import numpy as np
10
+
11
+ from contextdb.core.models import Edge, MemoryStatus
12
+ from contextdb.graphs.base import BaseGraph
13
+
14
+ if TYPE_CHECKING:
15
+ from contextdb.store.sqlite_store import SQLiteStore
16
+
17
+ _SCHEMA = """
18
+ CREATE TABLE IF NOT EXISTS semantic_edges (
19
+ source_id TEXT NOT NULL,
20
+ target_id TEXT NOT NULL,
21
+ weight REAL NOT NULL,
22
+ metadata TEXT DEFAULT '{}',
23
+ created_at TEXT NOT NULL,
24
+ PRIMARY KEY (source_id, target_id)
25
+ );
26
+ CREATE INDEX IF NOT EXISTS idx_semantic_source ON semantic_edges(source_id);
27
+ CREATE INDEX IF NOT EXISTS idx_semantic_target ON semantic_edges(target_id);
28
+ """
29
+
30
+
31
+ def _cosine(a: list[float], b: list[float]) -> float:
32
+ va = np.asarray(a, dtype=np.float32)
33
+ vb = np.asarray(b, dtype=np.float32)
34
+ na = float(np.linalg.norm(va))
35
+ nb = float(np.linalg.norm(vb))
36
+ if na == 0 or nb == 0:
37
+ return 0.0
38
+ return float(np.dot(va, vb) / (na * nb))
39
+
40
+
41
+ class SemanticGraph(BaseGraph):
42
+ """Edges represent cosine similarity above a threshold."""
43
+
44
+ def __init__(self, store: SQLiteStore, threshold: float = 0.6) -> None:
45
+ super().__init__(store)
46
+ self.threshold = threshold
47
+
48
+ async def initialize(self) -> None:
49
+ conn = self.store._require_conn()
50
+ await conn.executescript(_SCHEMA)
51
+ await conn.commit()
52
+
53
+ async def add_node(self, memory_id: str, data: dict[str, Any]) -> None:
54
+ embedding = data.get("embedding")
55
+ if embedding is None:
56
+ return
57
+ similar = await self.store.search_by_embedding(embedding, top_k=50)
58
+ for item in similar:
59
+ if item.id == memory_id or item.embedding is None:
60
+ continue
61
+ if item.status != MemoryStatus.ACTIVE:
62
+ continue
63
+ sim = _cosine(embedding, item.embedding)
64
+ if sim >= self.threshold:
65
+ await self.add_edge(
66
+ Edge(
67
+ source_id=memory_id,
68
+ target_id=item.id,
69
+ graph_type="semantic",
70
+ weight=sim,
71
+ )
72
+ )
73
+
74
+ async def add_edge(self, edge: Edge) -> None:
75
+ conn = self.store._require_conn()
76
+ await conn.execute(
77
+ "INSERT OR REPLACE INTO semantic_edges "
78
+ "(source_id, target_id, weight, metadata, created_at) VALUES (?,?,?,?,?)",
79
+ (
80
+ edge.source_id,
81
+ edge.target_id,
82
+ edge.weight,
83
+ json.dumps(edge.metadata),
84
+ edge.created_at.isoformat(),
85
+ ),
86
+ )
87
+ # Bidirectional for undirected similarity semantics.
88
+ await conn.execute(
89
+ "INSERT OR REPLACE INTO semantic_edges "
90
+ "(source_id, target_id, weight, metadata, created_at) VALUES (?,?,?,?,?)",
91
+ (
92
+ edge.target_id,
93
+ edge.source_id,
94
+ edge.weight,
95
+ json.dumps(edge.metadata),
96
+ edge.created_at.isoformat(),
97
+ ),
98
+ )
99
+ await conn.commit()
100
+
101
+ async def get_neighbors(
102
+ self,
103
+ memory_id: str,
104
+ depth: int = 1,
105
+ max_results: int = 20,
106
+ ) -> list[tuple[str, float]]:
107
+ conn = self.store._require_conn()
108
+ visited: dict[str, float] = {memory_id: 0.0}
109
+ frontier: list[tuple[str, float]] = [(memory_id, 1.0)]
110
+ for _ in range(depth):
111
+ next_frontier: list[tuple[str, float]] = []
112
+ for node, cum in frontier:
113
+ cursor = await conn.execute(
114
+ "SELECT target_id, weight FROM semantic_edges WHERE source_id = ?",
115
+ (node,),
116
+ )
117
+ rows = await cursor.fetchall()
118
+ for row in rows:
119
+ tid = row["target_id"]
120
+ w = cum * float(row["weight"])
121
+ if tid in visited and visited[tid] >= w:
122
+ continue
123
+ visited[tid] = w
124
+ next_frontier.append((tid, w))
125
+ frontier = next_frontier
126
+ visited.pop(memory_id, None)
127
+ ranked = sorted(visited.items(), key=lambda kv: kv[1], reverse=True)
128
+ return ranked[:max_results]
129
+
130
+ async def remove_node(self, memory_id: str) -> None:
131
+ conn = self.store._require_conn()
132
+ await conn.execute(
133
+ "DELETE FROM semantic_edges WHERE source_id = ? OR target_id = ?",
134
+ (memory_id, memory_id),
135
+ )
136
+ await conn.commit()
137
+
138
+ async def get_edges(self, memory_id: str) -> list[Edge]:
139
+ conn = self.store._require_conn()
140
+ cursor = await conn.execute(
141
+ "SELECT source_id, target_id, weight, metadata, created_at "
142
+ "FROM semantic_edges WHERE source_id = ?",
143
+ (memory_id,),
144
+ )
145
+ rows = await cursor.fetchall()
146
+ return [
147
+ Edge(
148
+ source_id=row["source_id"],
149
+ target_id=row["target_id"],
150
+ graph_type="semantic",
151
+ weight=float(row["weight"]),
152
+ metadata=json.loads(row["metadata"] or "{}"),
153
+ created_at=datetime.fromisoformat(row["created_at"]),
154
+ )
155
+ for row in rows
156
+ ]