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
|
@@ -0,0 +1,241 @@
|
|
|
1
|
+
"""Vector indices for fast nearest-neighbor search.
|
|
2
|
+
|
|
3
|
+
Two implementations ship with ContextDB:
|
|
4
|
+
|
|
5
|
+
* :class:`FAISSIndex` — uses ``faiss-cpu`` when installed; suitable for
|
|
6
|
+
collections up to tens of millions of vectors.
|
|
7
|
+
* :class:`NumpyIndex` — pure-numpy brute force; used as a fallback and for
|
|
8
|
+
test determinism. Fine for <10K vectors.
|
|
9
|
+
|
|
10
|
+
Callers should prefer :func:`get_vector_index`, which picks the best available
|
|
11
|
+
implementation at runtime.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import pickle
|
|
17
|
+
from abc import ABC, abstractmethod
|
|
18
|
+
from pathlib import Path
|
|
19
|
+
from typing import Any, cast
|
|
20
|
+
|
|
21
|
+
import numpy as np
|
|
22
|
+
from numpy.typing import NDArray
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _normalize(vectors: NDArray[np.float32]) -> NDArray[np.float32]:
|
|
26
|
+
"""L2-normalize rows. Zero vectors pass through unchanged."""
|
|
27
|
+
norms = np.linalg.norm(vectors, axis=1, keepdims=True)
|
|
28
|
+
norms = np.where(norms == 0, 1.0, norms)
|
|
29
|
+
return cast("NDArray[np.float32]", (vectors / norms).astype(np.float32))
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class VectorIndex(ABC):
|
|
33
|
+
"""Minimal vector index surface used by :class:`SQLiteStore`."""
|
|
34
|
+
|
|
35
|
+
@abstractmethod
|
|
36
|
+
def add(self, ids: list[str], embeddings: NDArray[np.float32]) -> None: ...
|
|
37
|
+
|
|
38
|
+
@abstractmethod
|
|
39
|
+
def search(self, query: NDArray[np.float32], top_k: int = 10) -> list[tuple[str, float]]: ...
|
|
40
|
+
|
|
41
|
+
@abstractmethod
|
|
42
|
+
def remove(self, ids: list[str]) -> None: ...
|
|
43
|
+
|
|
44
|
+
@abstractmethod
|
|
45
|
+
def save(self, path: str) -> None: ...
|
|
46
|
+
|
|
47
|
+
@abstractmethod
|
|
48
|
+
def load(self, path: str) -> None: ...
|
|
49
|
+
|
|
50
|
+
@abstractmethod
|
|
51
|
+
def __len__(self) -> int: ...
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class NumpyIndex(VectorIndex):
|
|
55
|
+
"""Brute-force cosine similarity over stacked float32 vectors."""
|
|
56
|
+
|
|
57
|
+
def __init__(self, dimension: int) -> None:
|
|
58
|
+
self.dimension = dimension
|
|
59
|
+
self._ids: list[str] = []
|
|
60
|
+
self._vectors: NDArray[np.float32] = np.zeros((0, dimension), dtype=np.float32)
|
|
61
|
+
|
|
62
|
+
def add(self, ids: list[str], embeddings: NDArray[np.float32]) -> None:
|
|
63
|
+
if not ids:
|
|
64
|
+
return
|
|
65
|
+
reshaped = np.asarray(embeddings, dtype=np.float32).reshape(len(ids), self.dimension)
|
|
66
|
+
normalized = _normalize(reshaped)
|
|
67
|
+
self._ids.extend(ids)
|
|
68
|
+
self._vectors = (
|
|
69
|
+
np.vstack([self._vectors, normalized]) if len(self._vectors) else normalized
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
def search(self, query: NDArray[np.float32], top_k: int = 10) -> list[tuple[str, float]]:
|
|
73
|
+
if len(self._ids) == 0:
|
|
74
|
+
return []
|
|
75
|
+
q = _normalize(np.asarray(query, dtype=np.float32).reshape(1, self.dimension))[0]
|
|
76
|
+
scores = self._vectors @ q
|
|
77
|
+
k = min(top_k, len(self._ids))
|
|
78
|
+
top_idx = np.argsort(-scores)[:k]
|
|
79
|
+
return [(self._ids[i], float(scores[i])) for i in top_idx]
|
|
80
|
+
|
|
81
|
+
def remove(self, ids: list[str]) -> None:
|
|
82
|
+
drop = set(ids)
|
|
83
|
+
keep = [i for i, mid in enumerate(self._ids) if mid not in drop]
|
|
84
|
+
self._ids = [self._ids[i] for i in keep]
|
|
85
|
+
self._vectors = self._vectors[keep] if keep else np.zeros(
|
|
86
|
+
(0, self.dimension), dtype=np.float32
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
def save(self, path: str) -> None:
|
|
90
|
+
payload = {
|
|
91
|
+
"dimension": self.dimension,
|
|
92
|
+
"ids": self._ids,
|
|
93
|
+
"vectors": self._vectors,
|
|
94
|
+
}
|
|
95
|
+
Path(path).write_bytes(pickle.dumps(payload))
|
|
96
|
+
|
|
97
|
+
def load(self, path: str) -> None:
|
|
98
|
+
payload = pickle.loads(Path(path).read_bytes())
|
|
99
|
+
self.dimension = int(payload["dimension"])
|
|
100
|
+
self._ids = list(payload["ids"])
|
|
101
|
+
self._vectors = np.asarray(payload["vectors"], dtype=np.float32)
|
|
102
|
+
|
|
103
|
+
def __len__(self) -> int:
|
|
104
|
+
return len(self._ids)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
class FAISSIndex(VectorIndex):
|
|
108
|
+
"""FAISS-backed index; requires the ``faiss-cpu`` optional dependency.
|
|
109
|
+
|
|
110
|
+
Removal is O(1): an id goes into :attr:`_removed_ids` and is filtered out
|
|
111
|
+
of every :meth:`search` call. The underlying FAISS index is rebuilt when
|
|
112
|
+
the tombstone set exceeds :attr:`_rebuild_threshold` (10% by default) so
|
|
113
|
+
search over-fetch stays bounded and index memory does not grow without
|
|
114
|
+
limit.
|
|
115
|
+
"""
|
|
116
|
+
|
|
117
|
+
_rebuild_threshold: float = 0.1
|
|
118
|
+
|
|
119
|
+
def __init__(self, dimension: int, index_type: str = "flat") -> None:
|
|
120
|
+
try:
|
|
121
|
+
import faiss
|
|
122
|
+
except ImportError as exc: # pragma: no cover - exercised only without faiss
|
|
123
|
+
raise RuntimeError(
|
|
124
|
+
"FAISS is not installed. Install with `pip install pycontextdb[faiss]` "
|
|
125
|
+
"or use NumpyIndex."
|
|
126
|
+
) from exc
|
|
127
|
+
self._faiss = faiss
|
|
128
|
+
self.dimension = dimension
|
|
129
|
+
self.index_type = index_type
|
|
130
|
+
self._index: Any = faiss.IndexFlatIP(dimension)
|
|
131
|
+
self._ids: list[str] = []
|
|
132
|
+
self._removed_ids: set[str] = set()
|
|
133
|
+
|
|
134
|
+
def add(self, ids: list[str], embeddings: NDArray[np.float32]) -> None:
|
|
135
|
+
if not ids:
|
|
136
|
+
return
|
|
137
|
+
reshaped = np.asarray(embeddings, dtype=np.float32).reshape(len(ids), self.dimension)
|
|
138
|
+
normalized = _normalize(reshaped)
|
|
139
|
+
self._index.add(normalized)
|
|
140
|
+
self._ids.extend(ids)
|
|
141
|
+
# Re-adding a previously tombstoned id clears the tombstone.
|
|
142
|
+
if self._removed_ids:
|
|
143
|
+
self._removed_ids.difference_update(ids)
|
|
144
|
+
|
|
145
|
+
def search(self, query: NDArray[np.float32], top_k: int = 10) -> list[tuple[str, float]]:
|
|
146
|
+
if not self._ids:
|
|
147
|
+
return []
|
|
148
|
+
# Over-fetch when tombstones are present so filtered-out hits do not
|
|
149
|
+
# starve the final top-k.
|
|
150
|
+
live = len(self._ids) - len(self._removed_ids)
|
|
151
|
+
if live <= 0:
|
|
152
|
+
return []
|
|
153
|
+
fetch = min(top_k + len(self._removed_ids), len(self._ids))
|
|
154
|
+
q = _normalize(np.asarray(query, dtype=np.float32).reshape(1, self.dimension))
|
|
155
|
+
scores, indices = self._index.search(q, fetch)
|
|
156
|
+
out: list[tuple[str, float]] = []
|
|
157
|
+
for score, idx in zip(scores[0], indices[0], strict=False):
|
|
158
|
+
if idx == -1:
|
|
159
|
+
continue
|
|
160
|
+
mid = self._ids[int(idx)]
|
|
161
|
+
if mid in self._removed_ids:
|
|
162
|
+
continue
|
|
163
|
+
out.append((mid, float(score)))
|
|
164
|
+
if len(out) >= top_k:
|
|
165
|
+
break
|
|
166
|
+
return out
|
|
167
|
+
|
|
168
|
+
def remove(self, ids: list[str]) -> None:
|
|
169
|
+
"""Tombstone the given ids. Amortized O(1) per id.
|
|
170
|
+
|
|
171
|
+
When tombstones exceed :attr:`_rebuild_threshold` of the live set,
|
|
172
|
+
:meth:`rebuild` is triggered to reclaim memory.
|
|
173
|
+
"""
|
|
174
|
+
if not ids:
|
|
175
|
+
return
|
|
176
|
+
present = {mid for mid in ids if mid in self._ids}
|
|
177
|
+
if not present:
|
|
178
|
+
return
|
|
179
|
+
self._removed_ids.update(present)
|
|
180
|
+
total = len(self._ids)
|
|
181
|
+
if total and len(self._removed_ids) / total > self._rebuild_threshold:
|
|
182
|
+
self.rebuild()
|
|
183
|
+
|
|
184
|
+
def rebuild(self) -> None:
|
|
185
|
+
"""Reconstruct the underlying FAISS index without the tombstoned ids.
|
|
186
|
+
|
|
187
|
+
Expensive (O(n) over live vectors) but amortized across removals.
|
|
188
|
+
"""
|
|
189
|
+
if not self._removed_ids:
|
|
190
|
+
return
|
|
191
|
+
keep_idx = [i for i, mid in enumerate(self._ids) if mid not in self._removed_ids]
|
|
192
|
+
new_ids = [self._ids[i] for i in keep_idx]
|
|
193
|
+
if keep_idx:
|
|
194
|
+
snapshot = self._vectors_snapshot()
|
|
195
|
+
vectors = np.stack([snapshot[i] for i in keep_idx], axis=0).astype(np.float32)
|
|
196
|
+
else:
|
|
197
|
+
vectors = np.zeros((0, self.dimension), dtype=np.float32)
|
|
198
|
+
self._index = self._faiss.IndexFlatIP(self.dimension)
|
|
199
|
+
self._ids = []
|
|
200
|
+
self._removed_ids.clear()
|
|
201
|
+
if len(new_ids):
|
|
202
|
+
self._index.add(vectors)
|
|
203
|
+
self._ids = new_ids
|
|
204
|
+
|
|
205
|
+
def _vectors_snapshot(self) -> NDArray[np.float32]:
|
|
206
|
+
# Reconstruct all current vectors from the FAISS index.
|
|
207
|
+
return np.stack(
|
|
208
|
+
[self._index.reconstruct(i) for i in range(self._index.ntotal)], axis=0
|
|
209
|
+
).astype(np.float32)
|
|
210
|
+
|
|
211
|
+
def save(self, path: str) -> None:
|
|
212
|
+
# Persisted state always reflects the post-rebuild, tombstone-free view
|
|
213
|
+
# so stale ids never leak into a reloaded index.
|
|
214
|
+
if self._removed_ids:
|
|
215
|
+
self.rebuild()
|
|
216
|
+
self._faiss.write_index(self._index, f"{path}.faiss")
|
|
217
|
+
Path(f"{path}.ids").write_bytes(pickle.dumps(self._ids))
|
|
218
|
+
|
|
219
|
+
def load(self, path: str) -> None:
|
|
220
|
+
self._index = self._faiss.read_index(f"{path}.faiss")
|
|
221
|
+
self._ids = pickle.loads(Path(f"{path}.ids").read_bytes())
|
|
222
|
+
self._removed_ids = set()
|
|
223
|
+
|
|
224
|
+
def __len__(self) -> int:
|
|
225
|
+
return len(self._ids) - len(self._removed_ids)
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def get_vector_index(dimension: int, prefer: str = "auto") -> VectorIndex:
|
|
229
|
+
"""Return the best available index for ``dimension``.
|
|
230
|
+
|
|
231
|
+
``prefer`` may be ``"faiss"``, ``"numpy"``, or ``"auto"`` (the default,
|
|
232
|
+
which uses FAISS if installed and falls back to NumPy otherwise).
|
|
233
|
+
"""
|
|
234
|
+
if prefer == "numpy":
|
|
235
|
+
return NumpyIndex(dimension)
|
|
236
|
+
if prefer == "faiss":
|
|
237
|
+
return FAISSIndex(dimension)
|
|
238
|
+
try:
|
|
239
|
+
return FAISSIndex(dimension)
|
|
240
|
+
except RuntimeError:
|
|
241
|
+
return NumpyIndex(dimension)
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
"""Utility adapters for embeddings and LLMs."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from contextdb.utils.embeddings import (
|
|
6
|
+
EmbeddingProvider,
|
|
7
|
+
MockEmbedding,
|
|
8
|
+
OpenAIEmbedding,
|
|
9
|
+
get_embedding_provider,
|
|
10
|
+
)
|
|
11
|
+
from contextdb.utils.llm import LLMProvider, MockLLM, OpenAILLM, get_llm_provider
|
|
12
|
+
|
|
13
|
+
__all__ = [
|
|
14
|
+
"EmbeddingProvider",
|
|
15
|
+
"LLMProvider",
|
|
16
|
+
"MockEmbedding",
|
|
17
|
+
"MockLLM",
|
|
18
|
+
"OpenAIEmbedding",
|
|
19
|
+
"OpenAILLM",
|
|
20
|
+
"get_embedding_provider",
|
|
21
|
+
"get_llm_provider",
|
|
22
|
+
]
|
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
"""Embedding providers.
|
|
2
|
+
|
|
3
|
+
ContextDB speaks a minimal :class:`EmbeddingProvider` protocol so swapping
|
|
4
|
+
backends (OpenAI ↔ local sentence-transformers ↔ deterministic mock) is a
|
|
5
|
+
one-line change. :func:`get_embedding_provider` is the factory most callers
|
|
6
|
+
want.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import hashlib
|
|
13
|
+
from abc import ABC, abstractmethod
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
import numpy as np
|
|
17
|
+
|
|
18
|
+
from contextdb.core.exceptions import ConfigError
|
|
19
|
+
|
|
20
|
+
# Known OpenAI embedding dimensions; update as new models ship.
|
|
21
|
+
_OPENAI_DIMS: dict[str, int] = {
|
|
22
|
+
"text-embedding-3-small": 1536,
|
|
23
|
+
"text-embedding-3-large": 3072,
|
|
24
|
+
"text-embedding-ada-002": 1536,
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class EmbeddingProvider(ABC):
|
|
29
|
+
"""Async embedding contract. Implementations must be batch-safe."""
|
|
30
|
+
|
|
31
|
+
@abstractmethod
|
|
32
|
+
async def embed(self, texts: list[str]) -> list[list[float]]: ...
|
|
33
|
+
|
|
34
|
+
@abstractmethod
|
|
35
|
+
def dimension(self) -> int: ...
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class OpenAIEmbedding(EmbeddingProvider):
|
|
39
|
+
"""OpenAI embedding API wrapper with exponential-backoff retry."""
|
|
40
|
+
|
|
41
|
+
def __init__(
|
|
42
|
+
self,
|
|
43
|
+
model: str = "text-embedding-3-small",
|
|
44
|
+
api_key: str | None = None,
|
|
45
|
+
max_retries: int = 3,
|
|
46
|
+
) -> None:
|
|
47
|
+
from openai import AsyncOpenAI
|
|
48
|
+
|
|
49
|
+
self.model = model
|
|
50
|
+
self.max_retries = max_retries
|
|
51
|
+
self._client = AsyncOpenAI(api_key=api_key)
|
|
52
|
+
self._dim = _OPENAI_DIMS.get(model, 1536)
|
|
53
|
+
|
|
54
|
+
async def embed(self, texts: list[str]) -> list[list[float]]:
|
|
55
|
+
if not texts:
|
|
56
|
+
return []
|
|
57
|
+
# OpenAI's per-call limit is 2048 inputs; chunk defensively.
|
|
58
|
+
chunks = [texts[i : i + 2048] for i in range(0, len(texts), 2048)]
|
|
59
|
+
out: list[list[float]] = []
|
|
60
|
+
for chunk in chunks:
|
|
61
|
+
out.extend(await self._embed_with_retry(chunk))
|
|
62
|
+
return out
|
|
63
|
+
|
|
64
|
+
async def _embed_with_retry(self, texts: list[str]) -> list[list[float]]:
|
|
65
|
+
delay = 1.0
|
|
66
|
+
for attempt in range(self.max_retries):
|
|
67
|
+
try:
|
|
68
|
+
response = await self._client.embeddings.create(
|
|
69
|
+
model=self.model, input=texts
|
|
70
|
+
)
|
|
71
|
+
return [d.embedding for d in response.data]
|
|
72
|
+
except Exception: # noqa: BLE001
|
|
73
|
+
if attempt == self.max_retries - 1:
|
|
74
|
+
raise
|
|
75
|
+
await asyncio.sleep(delay)
|
|
76
|
+
delay *= 2
|
|
77
|
+
return [] # pragma: no cover
|
|
78
|
+
|
|
79
|
+
def dimension(self) -> int:
|
|
80
|
+
return self._dim
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
class SentenceTransformerEmbedding(EmbeddingProvider):
|
|
84
|
+
"""Local-model embeddings via the optional ``sentence-transformers`` dep."""
|
|
85
|
+
|
|
86
|
+
def __init__(self, model_name: str = "all-MiniLM-L6-v2") -> None:
|
|
87
|
+
try:
|
|
88
|
+
from sentence_transformers import SentenceTransformer
|
|
89
|
+
except ImportError as exc: # pragma: no cover - optional dep
|
|
90
|
+
raise RuntimeError(
|
|
91
|
+
"sentence-transformers is not installed. "
|
|
92
|
+
"Install with `pip install pycontextdb[local]`."
|
|
93
|
+
) from exc
|
|
94
|
+
self._model = SentenceTransformer(model_name)
|
|
95
|
+
self._dim = int(self._model.get_sentence_embedding_dimension())
|
|
96
|
+
|
|
97
|
+
async def embed(self, texts: list[str]) -> list[list[float]]:
|
|
98
|
+
if not texts:
|
|
99
|
+
return []
|
|
100
|
+
loop = asyncio.get_running_loop()
|
|
101
|
+
|
|
102
|
+
def _run() -> list[list[float]]:
|
|
103
|
+
return [list(map(float, v)) for v in self._model.encode(texts)]
|
|
104
|
+
|
|
105
|
+
return await loop.run_in_executor(None, _run)
|
|
106
|
+
|
|
107
|
+
def dimension(self) -> int:
|
|
108
|
+
return self._dim
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
class MockEmbedding(EmbeddingProvider):
|
|
112
|
+
"""Deterministic pseudo-random embeddings for tests; no network calls."""
|
|
113
|
+
|
|
114
|
+
def __init__(self, dimension: int = 384) -> None:
|
|
115
|
+
self._dim = dimension
|
|
116
|
+
|
|
117
|
+
async def embed(self, texts: list[str]) -> list[list[float]]:
|
|
118
|
+
return [self._text_to_vector(t) for t in texts]
|
|
119
|
+
|
|
120
|
+
def _text_to_vector(self, text: str) -> list[float]:
|
|
121
|
+
digest = hashlib.sha256(text.encode("utf-8")).digest()
|
|
122
|
+
seed = int.from_bytes(digest[:4], "big")
|
|
123
|
+
rng = np.random.default_rng(seed)
|
|
124
|
+
vec = rng.standard_normal(self._dim).astype(np.float32)
|
|
125
|
+
# Encourage discrimination by mixing in per-word contribution.
|
|
126
|
+
for word in text.lower().split():
|
|
127
|
+
word_digest = hashlib.md5(word.encode("utf-8"), usedforsecurity=False).digest()
|
|
128
|
+
w_seed = int.from_bytes(word_digest[:4], "big")
|
|
129
|
+
word_rng = np.random.default_rng(w_seed)
|
|
130
|
+
vec += 0.5 * word_rng.standard_normal(self._dim).astype(np.float32)
|
|
131
|
+
norm = float(np.linalg.norm(vec))
|
|
132
|
+
if norm > 0:
|
|
133
|
+
vec = vec / norm
|
|
134
|
+
return [float(x) for x in vec]
|
|
135
|
+
|
|
136
|
+
def dimension(self) -> int:
|
|
137
|
+
return self._dim
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def get_embedding_provider(
|
|
141
|
+
model: str,
|
|
142
|
+
api_key: str | None = None,
|
|
143
|
+
**kwargs: Any,
|
|
144
|
+
) -> EmbeddingProvider:
|
|
145
|
+
"""Pick an embedding provider from a short model string.
|
|
146
|
+
|
|
147
|
+
``mock``/``test`` → :class:`MockEmbedding`; ``text-embedding-*`` or
|
|
148
|
+
``openai:*`` → :class:`OpenAIEmbedding`; anything else is routed to
|
|
149
|
+
:class:`SentenceTransformerEmbedding`.
|
|
150
|
+
"""
|
|
151
|
+
if model in {"mock", "test"}:
|
|
152
|
+
return MockEmbedding(dimension=kwargs.get("dimension", 384))
|
|
153
|
+
if model.startswith("text-embedding-") or model.startswith("openai:"):
|
|
154
|
+
resolved = model.replace("openai:", "", 1) if model.startswith("openai:") else model
|
|
155
|
+
return OpenAIEmbedding(model=resolved, api_key=api_key)
|
|
156
|
+
try:
|
|
157
|
+
return SentenceTransformerEmbedding(model_name=model)
|
|
158
|
+
except RuntimeError as exc:
|
|
159
|
+
raise ConfigError(str(exc)) from exc
|
contextdb/utils/llm.py
ADDED
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
"""LLM providers.
|
|
2
|
+
|
|
3
|
+
Minimal async contract over an LLM chat/completion, used by extraction,
|
|
4
|
+
compression, causal inference, and RL-as-policy pathways. Structured output is
|
|
5
|
+
approximated by asking the model for JSON and validating at the call site.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import asyncio
|
|
11
|
+
from abc import ABC, abstractmethod
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
from contextdb.core.exceptions import ConfigError
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class LLMProvider(ABC):
|
|
18
|
+
"""Common interface for text-in / text-out LLMs."""
|
|
19
|
+
|
|
20
|
+
@abstractmethod
|
|
21
|
+
async def generate(
|
|
22
|
+
self,
|
|
23
|
+
prompt: str,
|
|
24
|
+
system: str = "",
|
|
25
|
+
temperature: float = 0.0,
|
|
26
|
+
max_tokens: int = 1000,
|
|
27
|
+
response_format: type | None = None,
|
|
28
|
+
) -> str: ...
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class OpenAILLM(LLMProvider):
|
|
32
|
+
"""OpenAI chat-completions wrapper with retry."""
|
|
33
|
+
|
|
34
|
+
def __init__(
|
|
35
|
+
self,
|
|
36
|
+
model: str = "gpt-4o-mini",
|
|
37
|
+
api_key: str | None = None,
|
|
38
|
+
max_retries: int = 3,
|
|
39
|
+
) -> None:
|
|
40
|
+
from openai import AsyncOpenAI
|
|
41
|
+
|
|
42
|
+
self.model = model
|
|
43
|
+
self.max_retries = max_retries
|
|
44
|
+
self._client = AsyncOpenAI(api_key=api_key)
|
|
45
|
+
|
|
46
|
+
async def generate(
|
|
47
|
+
self,
|
|
48
|
+
prompt: str,
|
|
49
|
+
system: str = "",
|
|
50
|
+
temperature: float = 0.0,
|
|
51
|
+
max_tokens: int = 1000,
|
|
52
|
+
response_format: type | None = None,
|
|
53
|
+
) -> str:
|
|
54
|
+
messages: list[dict[str, str]] = []
|
|
55
|
+
if system:
|
|
56
|
+
messages.append({"role": "system", "content": system})
|
|
57
|
+
messages.append({"role": "user", "content": prompt})
|
|
58
|
+
extra: dict[str, Any] = {}
|
|
59
|
+
if response_format is not None:
|
|
60
|
+
# Ask for JSON; callers validate against their Pydantic model.
|
|
61
|
+
extra["response_format"] = {"type": "json_object"}
|
|
62
|
+
|
|
63
|
+
delay = 1.0
|
|
64
|
+
for attempt in range(self.max_retries):
|
|
65
|
+
try:
|
|
66
|
+
response = await self._client.chat.completions.create(
|
|
67
|
+
model=self.model,
|
|
68
|
+
messages=messages, # type: ignore[arg-type]
|
|
69
|
+
temperature=temperature,
|
|
70
|
+
max_tokens=max_tokens,
|
|
71
|
+
**extra,
|
|
72
|
+
)
|
|
73
|
+
content = response.choices[0].message.content or ""
|
|
74
|
+
return content
|
|
75
|
+
except Exception: # noqa: BLE001
|
|
76
|
+
if attempt == self.max_retries - 1:
|
|
77
|
+
raise
|
|
78
|
+
await asyncio.sleep(delay)
|
|
79
|
+
delay *= 2
|
|
80
|
+
return "" # pragma: no cover
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
class MockLLM(LLMProvider):
|
|
84
|
+
"""Scriptable LLM for tests.
|
|
85
|
+
|
|
86
|
+
``responses`` maps substring keys to the string that should be returned
|
|
87
|
+
when that substring appears anywhere in the prompt. All calls are logged
|
|
88
|
+
on :attr:`calls` so tests can assert on prompt content.
|
|
89
|
+
"""
|
|
90
|
+
|
|
91
|
+
def __init__(
|
|
92
|
+
self,
|
|
93
|
+
responses: dict[str, str] | None = None,
|
|
94
|
+
default: str = '{"facts": [], "entities": []}',
|
|
95
|
+
) -> None:
|
|
96
|
+
self.responses = responses or {}
|
|
97
|
+
self.default = default
|
|
98
|
+
self.calls: list[dict[str, Any]] = []
|
|
99
|
+
|
|
100
|
+
async def generate(
|
|
101
|
+
self,
|
|
102
|
+
prompt: str,
|
|
103
|
+
system: str = "",
|
|
104
|
+
temperature: float = 0.0,
|
|
105
|
+
max_tokens: int = 1000,
|
|
106
|
+
response_format: type | None = None,
|
|
107
|
+
) -> str:
|
|
108
|
+
self.calls.append(
|
|
109
|
+
{
|
|
110
|
+
"prompt": prompt,
|
|
111
|
+
"system": system,
|
|
112
|
+
"temperature": temperature,
|
|
113
|
+
"max_tokens": max_tokens,
|
|
114
|
+
}
|
|
115
|
+
)
|
|
116
|
+
for key, response in self.responses.items():
|
|
117
|
+
if key in prompt:
|
|
118
|
+
return response
|
|
119
|
+
return self.default
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def get_llm_provider(
|
|
123
|
+
model: str,
|
|
124
|
+
api_key: str | None = None,
|
|
125
|
+
**kwargs: Any,
|
|
126
|
+
) -> LLMProvider:
|
|
127
|
+
"""Route a model name to its provider.
|
|
128
|
+
|
|
129
|
+
``mock``/``test`` → :class:`MockLLM`; ``gpt-*`` or ``o1-*`` →
|
|
130
|
+
:class:`OpenAILLM`. Anything else raises :class:`ConfigError`.
|
|
131
|
+
"""
|
|
132
|
+
if model in {"mock", "test"}:
|
|
133
|
+
return MockLLM(**kwargs)
|
|
134
|
+
if model.startswith("gpt-") or model.startswith("o1-") or model.startswith("openai:"):
|
|
135
|
+
resolved = model.replace("openai:", "", 1) if model.startswith("openai:") else model
|
|
136
|
+
return OpenAILLM(model=resolved, api_key=api_key)
|
|
137
|
+
raise ConfigError(
|
|
138
|
+
f"Unknown LLM model '{model}'. Use 'mock' for tests or 'gpt-*'/'o1-*' for OpenAI."
|
|
139
|
+
)
|