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.
- contextdb/__init__.py +72 -0
- contextdb/agents/__init__.py +8 -0
- contextdb/agents/memory_bus.py +80 -0
- contextdb/agents/rl_manager.py +89 -0
- contextdb/cli.py +107 -0
- contextdb/client.py +516 -0
- contextdb/core/__init__.py +41 -0
- contextdb/core/config.py +89 -0
- contextdb/core/exceptions.py +29 -0
- contextdb/core/models.py +151 -0
- contextdb/dynamics/__init__.py +25 -0
- contextdb/dynamics/evolution.py +168 -0
- contextdb/dynamics/formation.py +193 -0
- contextdb/dynamics/retrieval.py +130 -0
- contextdb/graphs/__init__.py +17 -0
- contextdb/graphs/base.py +46 -0
- contextdb/graphs/causal.py +224 -0
- contextdb/graphs/entity.py +251 -0
- contextdb/graphs/semantic.py +156 -0
- contextdb/graphs/temporal.py +173 -0
- contextdb/integrations/__init__.py +10 -0
- contextdb/integrations/autogen.py +39 -0
- contextdb/integrations/crewai.py +41 -0
- contextdb/integrations/langchain.py +132 -0
- contextdb/integrations/openai_tools.py +124 -0
- contextdb/memory/__init__.py +9 -0
- contextdb/memory/experiential.py +102 -0
- contextdb/memory/factual.py +58 -0
- contextdb/memory/working.py +90 -0
- contextdb/privacy/__init__.py +9 -0
- contextdb/privacy/audit.py +199 -0
- contextdb/privacy/pii_detector.py +173 -0
- contextdb/privacy/retention.py +99 -0
- contextdb/py.typed +0 -0
- contextdb/store/__init__.py +15 -0
- contextdb/store/base.py +67 -0
- contextdb/store/sqlite_store.py +517 -0
- contextdb/store/vector_index.py +241 -0
- contextdb/utils/__init__.py +22 -0
- contextdb/utils/embeddings.py +159 -0
- contextdb/utils/llm.py +139 -0
- contextdb/utils/migrations.py +159 -0
- pycontextdb-0.1.0.dist-info/METADATA +589 -0
- pycontextdb-0.1.0.dist-info/RECORD +47 -0
- pycontextdb-0.1.0.dist-info/WHEEL +4 -0
- pycontextdb-0.1.0.dist-info/entry_points.txt +2 -0
- pycontextdb-0.1.0.dist-info/licenses/LICENSE +190 -0
contextdb/graphs/base.py
ADDED
|
@@ -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
|
+
]
|