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.
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,7 @@
1
+ from .provider import EmbeddingProvider, FunctionEmbeddingProvider, LocalEmbeddingProvider, _cosine
2
+ from .store import EmbeddingStore
3
+
4
+ __all__ = [
5
+ "EmbeddingProvider", "FunctionEmbeddingProvider", "LocalEmbeddingProvider",
6
+ "_cosine", "EmbeddingStore",
7
+ ]
@@ -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]