agentdatabase 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.
- agentdatabase-0.1.0.dist-info/METADATA +847 -0
- agentdatabase-0.1.0.dist-info/RECORD +35 -0
- agentdatabase-0.1.0.dist-info/WHEEL +5 -0
- agentdatabase-0.1.0.dist-info/entry_points.txt +2 -0
- agentdatabase-0.1.0.dist-info/licenses/LICENSE +651 -0
- agentdatabase-0.1.0.dist-info/top_level.txt +1 -0
- agentdb/__init__.py +5 -0
- agentdb/adapters/claude_agent_sdk.py +831 -0
- agentdb/adapters/hermes.py +247 -0
- agentdb/backend.py +75 -0
- agentdb/core/__init__.py +19 -0
- agentdb/core/directory_tracking.py +59 -0
- agentdb/core/file_integrity.py +79 -0
- agentdb/core/models.py +90 -0
- agentdb/core/profiles.py +116 -0
- agentdb/core/store.py +373 -0
- agentdb/core/system.py +86 -0
- agentdb/embeddings/__init__.py +7 -0
- agentdb/embeddings/provider.py +99 -0
- agentdb/embeddings/store.py +199 -0
- agentdb/embeddings/text.py +20 -0
- agentdb/gateway/__init__.py +189 -0
- agentdb/gateway/adapter.py +58 -0
- agentdb/governance/__init__.py +3 -0
- agentdb/governance/conflict_detector.py +103 -0
- agentdb/governance/lifecycle_manager.py +226 -0
- agentdb/governance/permission_router.py +131 -0
- agentdb/interface/__init__.py +25 -0
- agentdb/interface/client.py +1108 -0
- agentdb/interface/mcp_server.py +96 -0
- agentdb/retrieval/__init__.py +3 -0
- agentdb/retrieval/algorithm.py +207 -0
- agentdb/skills/__init__.py +3 -0
- agentdb/skills/skill_store.py +351 -0
- agentdb/testing.py +68 -0
agentdb/core/store.py
ADDED
|
@@ -0,0 +1,373 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import sqlite3
|
|
5
|
+
from contextlib import contextmanager
|
|
6
|
+
from datetime import datetime, timezone
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Any, Optional
|
|
9
|
+
|
|
10
|
+
from .models import ConflictRecord, MemoryRecord
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
# ---------------------------------------------------------------------------
|
|
14
|
+
# SQLite schema
|
|
15
|
+
# ---------------------------------------------------------------------------
|
|
16
|
+
|
|
17
|
+
_CREATE_MEMORIES = """
|
|
18
|
+
CREATE TABLE IF NOT EXISTS memories (
|
|
19
|
+
id TEXT PRIMARY KEY,
|
|
20
|
+
key TEXT NOT NULL,
|
|
21
|
+
value TEXT NOT NULL,
|
|
22
|
+
origin TEXT NOT NULL,
|
|
23
|
+
trust_level TEXT NOT NULL,
|
|
24
|
+
agent_id TEXT NOT NULL,
|
|
25
|
+
confidence REAL NOT NULL DEFAULT 1.0,
|
|
26
|
+
version INTEGER NOT NULL DEFAULT 1,
|
|
27
|
+
archived INTEGER NOT NULL DEFAULT 0,
|
|
28
|
+
retrieval_count INTEGER NOT NULL DEFAULT 0,
|
|
29
|
+
created_at TEXT NOT NULL,
|
|
30
|
+
entities TEXT NOT NULL DEFAULT '[]',
|
|
31
|
+
last_retrieved TEXT,
|
|
32
|
+
conflict_id TEXT,
|
|
33
|
+
expiry_time TEXT,
|
|
34
|
+
inferred INTEGER NOT NULL DEFAULT 0,
|
|
35
|
+
working_dirs TEXT NOT NULL DEFAULT '[]'
|
|
36
|
+
)
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
_MEMORIES_MIGRATION_COLUMNS = (
|
|
40
|
+
("last_retrieved", "TEXT"),
|
|
41
|
+
("conflict_id", "TEXT"),
|
|
42
|
+
("expiry_time", "TEXT"),
|
|
43
|
+
("inferred", "INTEGER NOT NULL DEFAULT 0"),
|
|
44
|
+
("working_dirs", "TEXT NOT NULL DEFAULT '[]'"),
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
_CREATE_CONFLICTS = """
|
|
48
|
+
CREATE TABLE IF NOT EXISTS conflicts (
|
|
49
|
+
id TEXT PRIMARY KEY,
|
|
50
|
+
record_a_id TEXT NOT NULL,
|
|
51
|
+
record_b_id TEXT NOT NULL,
|
|
52
|
+
conflict_type TEXT NOT NULL,
|
|
53
|
+
resolved INTEGER NOT NULL DEFAULT 0,
|
|
54
|
+
auto_resolved INTEGER NOT NULL DEFAULT 0,
|
|
55
|
+
created_at TEXT NOT NULL
|
|
56
|
+
)
|
|
57
|
+
"""
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def _utcnow() -> datetime:
|
|
61
|
+
return datetime.now(timezone.utc)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
# ---------------------------------------------------------------------------
|
|
65
|
+
# In-memory store backed by SQLite
|
|
66
|
+
# ---------------------------------------------------------------------------
|
|
67
|
+
|
|
68
|
+
class MemoryStore:
|
|
69
|
+
"""SQLite-backed store for MemoryRecord objects."""
|
|
70
|
+
|
|
71
|
+
def __init__(self, db_path: Path) -> None:
|
|
72
|
+
self._db_path = db_path
|
|
73
|
+
self._by_id: dict[str, MemoryRecord] = {}
|
|
74
|
+
db_path.parent.mkdir(parents=True, exist_ok=True)
|
|
75
|
+
with self._connect() as conn:
|
|
76
|
+
conn.execute(_CREATE_MEMORIES)
|
|
77
|
+
conn.execute(_CREATE_CONFLICTS)
|
|
78
|
+
for column, coltype in _MEMORIES_MIGRATION_COLUMNS:
|
|
79
|
+
try:
|
|
80
|
+
conn.execute(f"ALTER TABLE memories ADD COLUMN {column} {coltype}")
|
|
81
|
+
except sqlite3.OperationalError:
|
|
82
|
+
pass # column already exists (pre-existing database)
|
|
83
|
+
# Load existing records into memory
|
|
84
|
+
self._reload()
|
|
85
|
+
|
|
86
|
+
@contextmanager
|
|
87
|
+
def _connect(self):
|
|
88
|
+
conn = sqlite3.connect(str(self._db_path))
|
|
89
|
+
conn.row_factory = sqlite3.Row
|
|
90
|
+
conn.execute("PRAGMA journal_mode=WAL")
|
|
91
|
+
try:
|
|
92
|
+
yield conn
|
|
93
|
+
conn.commit()
|
|
94
|
+
except Exception:
|
|
95
|
+
conn.rollback()
|
|
96
|
+
raise
|
|
97
|
+
finally:
|
|
98
|
+
conn.close()
|
|
99
|
+
|
|
100
|
+
def _reload(self) -> None:
|
|
101
|
+
with self._connect() as conn:
|
|
102
|
+
rows = conn.execute("SELECT * FROM memories").fetchall()
|
|
103
|
+
for row in rows:
|
|
104
|
+
record = self._row_to_record(row)
|
|
105
|
+
self._by_id[record.id] = record
|
|
106
|
+
|
|
107
|
+
def _row_to_record(self, row: sqlite3.Row) -> MemoryRecord:
|
|
108
|
+
return MemoryRecord(
|
|
109
|
+
id=row["id"],
|
|
110
|
+
key=row["key"],
|
|
111
|
+
value=json.loads(row["value"]),
|
|
112
|
+
origin=row["origin"],
|
|
113
|
+
trust_level=row["trust_level"],
|
|
114
|
+
agent_id=row["agent_id"],
|
|
115
|
+
confidence=row["confidence"],
|
|
116
|
+
version=row["version"],
|
|
117
|
+
archived=bool(row["archived"]),
|
|
118
|
+
retrieval_count=row["retrieval_count"],
|
|
119
|
+
created_at=datetime.fromisoformat(row["created_at"]),
|
|
120
|
+
entities=json.loads(row["entities"]),
|
|
121
|
+
last_retrieved=(
|
|
122
|
+
datetime.fromisoformat(row["last_retrieved"])
|
|
123
|
+
if row["last_retrieved"] else None
|
|
124
|
+
),
|
|
125
|
+
conflict_id=row["conflict_id"],
|
|
126
|
+
expiry_time=(
|
|
127
|
+
datetime.fromisoformat(row["expiry_time"])
|
|
128
|
+
if row["expiry_time"] else None
|
|
129
|
+
),
|
|
130
|
+
inferred=bool(row["inferred"]) if row["inferred"] is not None else False,
|
|
131
|
+
working_dirs=json.loads(row["working_dirs"] if "working_dirs" in row.keys() else "[]"),
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
def write(self, record: MemoryRecord) -> MemoryRecord:
|
|
135
|
+
with self._connect() as conn:
|
|
136
|
+
conn.execute(
|
|
137
|
+
"""
|
|
138
|
+
INSERT OR REPLACE INTO memories
|
|
139
|
+
(id, key, value, origin, trust_level, agent_id, confidence,
|
|
140
|
+
version, archived, retrieval_count, created_at, entities,
|
|
141
|
+
last_retrieved, conflict_id, expiry_time, inferred, working_dirs)
|
|
142
|
+
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
143
|
+
""",
|
|
144
|
+
(
|
|
145
|
+
record.id,
|
|
146
|
+
record.key,
|
|
147
|
+
json.dumps(record.value),
|
|
148
|
+
str(record.origin),
|
|
149
|
+
str(record.trust_level),
|
|
150
|
+
record.agent_id,
|
|
151
|
+
record.confidence,
|
|
152
|
+
record.version,
|
|
153
|
+
int(record.archived),
|
|
154
|
+
record.retrieval_count,
|
|
155
|
+
record.created_at.isoformat(),
|
|
156
|
+
json.dumps(record.entities),
|
|
157
|
+
record.last_retrieved.isoformat() if record.last_retrieved else None,
|
|
158
|
+
record.conflict_id,
|
|
159
|
+
record.expiry_time.isoformat() if record.expiry_time else None,
|
|
160
|
+
int(record.inferred),
|
|
161
|
+
json.dumps(record.working_dirs),
|
|
162
|
+
),
|
|
163
|
+
)
|
|
164
|
+
self._by_id[record.id] = record
|
|
165
|
+
return record
|
|
166
|
+
|
|
167
|
+
def get_by_key(self, key: str) -> Optional[MemoryRecord]:
|
|
168
|
+
for record in self._by_id.values():
|
|
169
|
+
if record.key == key and not record.archived:
|
|
170
|
+
return record
|
|
171
|
+
return None
|
|
172
|
+
|
|
173
|
+
def get_by_id(self, record_id: str) -> Optional[MemoryRecord]:
|
|
174
|
+
return self._by_id.get(record_id)
|
|
175
|
+
|
|
176
|
+
def delete_by_key(self, key: str) -> None:
|
|
177
|
+
record = self.get_by_key(key)
|
|
178
|
+
if record is None:
|
|
179
|
+
return
|
|
180
|
+
record.archived = True
|
|
181
|
+
self._persist_field(record.id, "archived", 1)
|
|
182
|
+
|
|
183
|
+
def _persist_field(self, record_id: str, field_name: str, value: Any) -> None:
|
|
184
|
+
with self._connect() as conn:
|
|
185
|
+
conn.execute(
|
|
186
|
+
f"UPDATE memories SET {field_name} = ? WHERE id = ?",
|
|
187
|
+
(value, record_id),
|
|
188
|
+
)
|
|
189
|
+
|
|
190
|
+
def increment_retrieval_count(self, record_id: str) -> None:
|
|
191
|
+
record = self._by_id.get(record_id)
|
|
192
|
+
if record:
|
|
193
|
+
record.retrieval_count += 1
|
|
194
|
+
self._persist_field(record_id, "retrieval_count", record.retrieval_count)
|
|
195
|
+
record.last_retrieved = _utcnow()
|
|
196
|
+
self._persist_field(record_id, "last_retrieved", record.last_retrieved.isoformat())
|
|
197
|
+
|
|
198
|
+
def update_confidence(self, record_id: str, new_confidence: float) -> None:
|
|
199
|
+
record = self._by_id.get(record_id)
|
|
200
|
+
if record:
|
|
201
|
+
record.confidence = round(new_confidence, 10)
|
|
202
|
+
self._persist_field(record_id, "confidence", record.confidence)
|
|
203
|
+
|
|
204
|
+
def archive_record(self, record_id: str) -> None:
|
|
205
|
+
"""Mark a record as archived by ID."""
|
|
206
|
+
record = self._by_id.get(record_id)
|
|
207
|
+
if record:
|
|
208
|
+
record.archived = True
|
|
209
|
+
self._persist_field(record_id, "archived", 1)
|
|
210
|
+
|
|
211
|
+
def list_all(self) -> list[MemoryRecord]:
|
|
212
|
+
return [r for r in self._by_id.values() if not r.archived]
|
|
213
|
+
|
|
214
|
+
def list_by_agent(self, agent_id: str) -> list[MemoryRecord]:
|
|
215
|
+
return [r for r in self._by_id.values() if r.agent_id == agent_id and not r.archived]
|
|
216
|
+
|
|
217
|
+
def search(self, query: str, limit: int = 5, agent_id: Optional[str] = None) -> list[MemoryRecord]:
|
|
218
|
+
"""
|
|
219
|
+
Simple text search: find records whose key or value contains
|
|
220
|
+
any word from query (case-insensitive). agent_id, if given, is
|
|
221
|
+
applied during the scan (before the `limit` slice below) so an
|
|
222
|
+
isolated caller's own matches can't be truncated away by another
|
|
223
|
+
agent's larger matching set.
|
|
224
|
+
"""
|
|
225
|
+
words = [w.lower() for w in query.split() if w]
|
|
226
|
+
if not words:
|
|
227
|
+
return []
|
|
228
|
+
|
|
229
|
+
results: list[MemoryRecord] = []
|
|
230
|
+
for record in self._by_id.values():
|
|
231
|
+
if record.archived:
|
|
232
|
+
continue
|
|
233
|
+
if agent_id is not None and record.agent_id != agent_id:
|
|
234
|
+
continue
|
|
235
|
+
key_lower = record.key.lower()
|
|
236
|
+
value_str = (
|
|
237
|
+
json.dumps(record.value) if isinstance(record.value, (dict, list)) else str(record.value)
|
|
238
|
+
).lower()
|
|
239
|
+
for word in words:
|
|
240
|
+
if word in key_lower or word in value_str:
|
|
241
|
+
results.append(record)
|
|
242
|
+
break
|
|
243
|
+
|
|
244
|
+
return results[:limit]
|
|
245
|
+
|
|
246
|
+
def search_by_entities(self, entities: list[str], agent_id: Optional[str] = None) -> list[MemoryRecord]:
|
|
247
|
+
"""
|
|
248
|
+
Find records where record.entities has overlap with the given entities.
|
|
249
|
+
Fallback: search key and value for entity name substrings.
|
|
250
|
+
agent_id, if given, is applied inline in BOTH passes below — not as a
|
|
251
|
+
filter on the final chosen list — so an isolated caller's fallback
|
|
252
|
+
search is scoped to their own records rather than being gated by
|
|
253
|
+
whether OTHER agents had entity-tag hits.
|
|
254
|
+
"""
|
|
255
|
+
if not entities:
|
|
256
|
+
return []
|
|
257
|
+
|
|
258
|
+
entity_set = set(e.lower() for e in entities)
|
|
259
|
+
|
|
260
|
+
# First pass: exact entity match
|
|
261
|
+
matches: list[MemoryRecord] = []
|
|
262
|
+
for record in self._by_id.values():
|
|
263
|
+
if record.archived:
|
|
264
|
+
continue
|
|
265
|
+
if agent_id is not None and record.agent_id != agent_id:
|
|
266
|
+
continue
|
|
267
|
+
record_entities = set(e.lower() for e in (record.entities or []))
|
|
268
|
+
if record_entities & entity_set:
|
|
269
|
+
matches.append(record)
|
|
270
|
+
|
|
271
|
+
if matches:
|
|
272
|
+
return matches
|
|
273
|
+
|
|
274
|
+
# Fallback: substring search in key and value
|
|
275
|
+
fallback: list[MemoryRecord] = []
|
|
276
|
+
for record in self._by_id.values():
|
|
277
|
+
if record.archived:
|
|
278
|
+
continue
|
|
279
|
+
if agent_id is not None and record.agent_id != agent_id:
|
|
280
|
+
continue
|
|
281
|
+
key_lower = record.key.lower()
|
|
282
|
+
value_str = (
|
|
283
|
+
json.dumps(record.value) if isinstance(record.value, (dict, list)) else str(record.value)
|
|
284
|
+
).lower()
|
|
285
|
+
for entity in entities:
|
|
286
|
+
entity_lower = entity.lower()
|
|
287
|
+
if entity_lower in key_lower or entity_lower in value_str:
|
|
288
|
+
fallback.append(record)
|
|
289
|
+
break
|
|
290
|
+
|
|
291
|
+
return fallback
|
|
292
|
+
|
|
293
|
+
# ------------------------------------------------------------------
|
|
294
|
+
# Conflict storage
|
|
295
|
+
# ------------------------------------------------------------------
|
|
296
|
+
|
|
297
|
+
def store_conflict(self, conflict: ConflictRecord) -> None:
|
|
298
|
+
now = _utcnow().isoformat()
|
|
299
|
+
with self._connect() as conn:
|
|
300
|
+
conn.execute(
|
|
301
|
+
"""
|
|
302
|
+
INSERT OR REPLACE INTO conflicts
|
|
303
|
+
(id, record_a_id, record_b_id, conflict_type, resolved, auto_resolved, created_at)
|
|
304
|
+
VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
305
|
+
""",
|
|
306
|
+
(
|
|
307
|
+
conflict.id,
|
|
308
|
+
conflict.record_a_id,
|
|
309
|
+
conflict.record_b_id,
|
|
310
|
+
conflict.conflict_type,
|
|
311
|
+
int(conflict.resolved),
|
|
312
|
+
int(conflict.auto_resolved),
|
|
313
|
+
now,
|
|
314
|
+
),
|
|
315
|
+
)
|
|
316
|
+
|
|
317
|
+
def get_conflicts_for_records(self, record_ids: list[str]) -> list[ConflictRecord]:
|
|
318
|
+
if not record_ids:
|
|
319
|
+
return []
|
|
320
|
+
placeholders = ",".join("?" * len(record_ids))
|
|
321
|
+
with self._connect() as conn:
|
|
322
|
+
rows = conn.execute(
|
|
323
|
+
f"""
|
|
324
|
+
SELECT id, record_a_id, record_b_id, conflict_type, resolved, auto_resolved
|
|
325
|
+
FROM conflicts
|
|
326
|
+
WHERE record_a_id IN ({placeholders}) OR record_b_id IN ({placeholders})
|
|
327
|
+
""",
|
|
328
|
+
record_ids + record_ids,
|
|
329
|
+
).fetchall()
|
|
330
|
+
return [
|
|
331
|
+
ConflictRecord(
|
|
332
|
+
id=row["id"],
|
|
333
|
+
record_a_id=row["record_a_id"],
|
|
334
|
+
record_b_id=row["record_b_id"],
|
|
335
|
+
conflict_type=row["conflict_type"],
|
|
336
|
+
resolved=bool(row["resolved"]),
|
|
337
|
+
auto_resolved=bool(row["auto_resolved"]),
|
|
338
|
+
)
|
|
339
|
+
for row in rows
|
|
340
|
+
]
|
|
341
|
+
|
|
342
|
+
def list_conflicts(self, resolved: bool = False) -> list[ConflictRecord]:
|
|
343
|
+
"""
|
|
344
|
+
Return conflict records from the conflicts table.
|
|
345
|
+
|
|
346
|
+
resolved=False (default) -> pending only (resolved = 0).
|
|
347
|
+
resolved=True -> all conflicts (pending + resolved).
|
|
348
|
+
"""
|
|
349
|
+
query = (
|
|
350
|
+
"SELECT id, record_a_id, record_b_id, conflict_type, resolved, auto_resolved "
|
|
351
|
+
"FROM conflicts"
|
|
352
|
+
)
|
|
353
|
+
with self._connect() as conn:
|
|
354
|
+
if not resolved:
|
|
355
|
+
query += " WHERE resolved = 0"
|
|
356
|
+
rows = conn.execute(query).fetchall()
|
|
357
|
+
return [
|
|
358
|
+
ConflictRecord(
|
|
359
|
+
id=row["id"],
|
|
360
|
+
record_a_id=row["record_a_id"],
|
|
361
|
+
record_b_id=row["record_b_id"],
|
|
362
|
+
conflict_type=row["conflict_type"],
|
|
363
|
+
resolved=bool(row["resolved"]),
|
|
364
|
+
auto_resolved=bool(row["auto_resolved"]),
|
|
365
|
+
)
|
|
366
|
+
for row in rows
|
|
367
|
+
]
|
|
368
|
+
|
|
369
|
+
def update_conflict_id(self, record_id: str, conflict_id: str) -> None:
|
|
370
|
+
self._persist_field(record_id, "conflict_id", conflict_id)
|
|
371
|
+
record = self._by_id.get(record_id)
|
|
372
|
+
if record is not None:
|
|
373
|
+
record.conflict_id = conflict_id
|
agentdb/core/system.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import sqlite3
|
|
4
|
+
from contextlib import contextmanager
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class SystemStats:
|
|
9
|
+
def __init__(self, db_path: Path) -> None:
|
|
10
|
+
self._db_path = db_path
|
|
11
|
+
|
|
12
|
+
@contextmanager
|
|
13
|
+
def _connect(self):
|
|
14
|
+
conn = sqlite3.connect(str(self._db_path))
|
|
15
|
+
conn.row_factory = sqlite3.Row
|
|
16
|
+
conn.execute("PRAGMA journal_mode=WAL")
|
|
17
|
+
try:
|
|
18
|
+
yield conn
|
|
19
|
+
conn.commit()
|
|
20
|
+
except Exception:
|
|
21
|
+
conn.rollback()
|
|
22
|
+
raise
|
|
23
|
+
finally:
|
|
24
|
+
conn.close()
|
|
25
|
+
|
|
26
|
+
def agentdb_stats(self) -> dict:
|
|
27
|
+
with self._connect() as conn:
|
|
28
|
+
# Total memories (non-archived)
|
|
29
|
+
total = conn.execute(
|
|
30
|
+
"SELECT COUNT(*) FROM memories WHERE archived = 0"
|
|
31
|
+
).fetchone()[0]
|
|
32
|
+
|
|
33
|
+
# Average confidence
|
|
34
|
+
avg_conf_row = conn.execute(
|
|
35
|
+
"SELECT AVG(confidence) FROM memories WHERE archived = 0"
|
|
36
|
+
).fetchone()
|
|
37
|
+
avg_confidence = avg_conf_row[0] if avg_conf_row[0] is not None else 0.0
|
|
38
|
+
|
|
39
|
+
# Low confidence count (< 0.5)
|
|
40
|
+
low_conf = conn.execute(
|
|
41
|
+
"SELECT COUNT(*) FROM memories WHERE archived = 0 AND confidence < 0.5"
|
|
42
|
+
).fetchone()[0]
|
|
43
|
+
|
|
44
|
+
# Retrieval count total (sum of retrieval_count across all memories)
|
|
45
|
+
retrieval_total_row = conn.execute(
|
|
46
|
+
"SELECT SUM(retrieval_count) FROM memories WHERE archived = 0"
|
|
47
|
+
).fetchone()
|
|
48
|
+
retrieval_total = retrieval_total_row[0] if retrieval_total_row[0] is not None else 0
|
|
49
|
+
|
|
50
|
+
# Conflict pending count - from audit log
|
|
51
|
+
conflict_pending = 0 # We don't track ConflictRecord table; use 0
|
|
52
|
+
|
|
53
|
+
# Outcome positive rate
|
|
54
|
+
outcome_rate = None
|
|
55
|
+
try:
|
|
56
|
+
outcomes_rows = conn.execute(
|
|
57
|
+
"SELECT outcome_value FROM outcomes"
|
|
58
|
+
).fetchall()
|
|
59
|
+
if len(outcomes_rows) >= 10:
|
|
60
|
+
positive = sum(1 for r in outcomes_rows if r["outcome_value"] >= 0.5)
|
|
61
|
+
outcome_rate = positive / len(outcomes_rows)
|
|
62
|
+
except Exception:
|
|
63
|
+
pass
|
|
64
|
+
|
|
65
|
+
# External conflict rate
|
|
66
|
+
external_conflict_rate = None
|
|
67
|
+
try:
|
|
68
|
+
ext_rows = conn.execute(
|
|
69
|
+
"SELECT had_external_context, had_conflict FROM external_retrieves"
|
|
70
|
+
).fetchall()
|
|
71
|
+
external_retrieves = [r for r in ext_rows if r["had_external_context"]]
|
|
72
|
+
if external_retrieves:
|
|
73
|
+
conflicts = sum(1 for r in external_retrieves if r["had_conflict"])
|
|
74
|
+
external_conflict_rate = conflicts / len(external_retrieves)
|
|
75
|
+
except Exception:
|
|
76
|
+
pass
|
|
77
|
+
|
|
78
|
+
return {
|
|
79
|
+
"total_memories": total,
|
|
80
|
+
"avg_confidence": avg_confidence or 0.0,
|
|
81
|
+
"low_confidence_count": low_conf,
|
|
82
|
+
"conflict_pending_count": conflict_pending,
|
|
83
|
+
"outcome_positive_rate": outcome_rate,
|
|
84
|
+
"retrieval_count_total": retrieval_total,
|
|
85
|
+
"external_conflict_rate": external_conflict_rate,
|
|
86
|
+
}
|
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
import threading
|
|
5
|
+
from typing import Callable, Protocol, runtime_checkable
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@runtime_checkable
|
|
9
|
+
class EmbeddingProvider(Protocol):
|
|
10
|
+
@property
|
|
11
|
+
def model_name(self) -> str: ...
|
|
12
|
+
@property
|
|
13
|
+
def dimensions(self) -> int: ...
|
|
14
|
+
def embed(self, text: str) -> list[float]: ...
|
|
15
|
+
def embed_batch(self, texts: list[str]) -> list[list[float]]: ...
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _cosine(a: list[float], b: list[float]) -> float:
|
|
19
|
+
"""Cosine similarity between two vectors. Safe for zero vectors."""
|
|
20
|
+
dot = sum(x * y for x, y in zip(a, b))
|
|
21
|
+
norm_a = math.sqrt(sum(x * x for x in a))
|
|
22
|
+
norm_b = math.sqrt(sum(x * x for x in b))
|
|
23
|
+
if norm_a == 0.0 or norm_b == 0.0:
|
|
24
|
+
return 0.0
|
|
25
|
+
return dot / (norm_a * norm_b)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class FunctionEmbeddingProvider:
|
|
29
|
+
def __init__(self, fn: Callable[[str], list[float]], model_name: str = "custom") -> None:
|
|
30
|
+
self._fn = fn
|
|
31
|
+
self._model_name = model_name
|
|
32
|
+
self._dimensions: int | None = None
|
|
33
|
+
|
|
34
|
+
@property
|
|
35
|
+
def model_name(self) -> str:
|
|
36
|
+
return self._model_name
|
|
37
|
+
|
|
38
|
+
@property
|
|
39
|
+
def dimensions(self) -> int:
|
|
40
|
+
if self._dimensions is None:
|
|
41
|
+
try:
|
|
42
|
+
probe = self._fn("probe")
|
|
43
|
+
self._dimensions = len(probe)
|
|
44
|
+
except Exception:
|
|
45
|
+
self._dimensions = 0
|
|
46
|
+
return self._dimensions
|
|
47
|
+
|
|
48
|
+
def embed(self, text: str) -> list[float]:
|
|
49
|
+
return self._fn(text)
|
|
50
|
+
|
|
51
|
+
def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
|
52
|
+
return [self._fn(t) for t in texts]
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class LocalEmbeddingProvider:
|
|
56
|
+
_MODEL_NAME = "sentence-transformers/all-MiniLM-L6-v2"
|
|
57
|
+
_model = None
|
|
58
|
+
_lock = threading.Lock()
|
|
59
|
+
|
|
60
|
+
@property
|
|
61
|
+
def model_name(self) -> str:
|
|
62
|
+
return self._MODEL_NAME
|
|
63
|
+
|
|
64
|
+
@property
|
|
65
|
+
def dimensions(self) -> int:
|
|
66
|
+
return 384
|
|
67
|
+
|
|
68
|
+
def _load_model(self) -> None:
|
|
69
|
+
if LocalEmbeddingProvider._model is None:
|
|
70
|
+
try:
|
|
71
|
+
from sentence_transformers import SentenceTransformer # type: ignore
|
|
72
|
+
except ImportError as e:
|
|
73
|
+
raise RuntimeError(
|
|
74
|
+
"sentence-transformers is required for LocalEmbeddingProvider. "
|
|
75
|
+
"Install it with: pip install 'agentdb[local-embed]'"
|
|
76
|
+
) from e
|
|
77
|
+
import logging
|
|
78
|
+
logging.getLogger("agentdb").info(
|
|
79
|
+
"agentdb: loading local embedding model (first use)"
|
|
80
|
+
)
|
|
81
|
+
LocalEmbeddingProvider._model = SentenceTransformer(self._MODEL_NAME)
|
|
82
|
+
|
|
83
|
+
def embed(self, text: str) -> list[float]:
|
|
84
|
+
with LocalEmbeddingProvider._lock:
|
|
85
|
+
self._load_model()
|
|
86
|
+
vecs = LocalEmbeddingProvider._model.encode(
|
|
87
|
+
[text], normalize_embeddings=True
|
|
88
|
+
)
|
|
89
|
+
return vecs[0].tolist()
|
|
90
|
+
|
|
91
|
+
def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
|
92
|
+
if not texts:
|
|
93
|
+
return []
|
|
94
|
+
with LocalEmbeddingProvider._lock:
|
|
95
|
+
self._load_model()
|
|
96
|
+
vecs = LocalEmbeddingProvider._model.encode(
|
|
97
|
+
texts, normalize_embeddings=True
|
|
98
|
+
)
|
|
99
|
+
return [v.tolist() for v in vecs]
|