pycontextdb 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (47) hide show
  1. contextdb/__init__.py +72 -0
  2. contextdb/agents/__init__.py +8 -0
  3. contextdb/agents/memory_bus.py +80 -0
  4. contextdb/agents/rl_manager.py +89 -0
  5. contextdb/cli.py +107 -0
  6. contextdb/client.py +516 -0
  7. contextdb/core/__init__.py +41 -0
  8. contextdb/core/config.py +89 -0
  9. contextdb/core/exceptions.py +29 -0
  10. contextdb/core/models.py +151 -0
  11. contextdb/dynamics/__init__.py +25 -0
  12. contextdb/dynamics/evolution.py +168 -0
  13. contextdb/dynamics/formation.py +193 -0
  14. contextdb/dynamics/retrieval.py +130 -0
  15. contextdb/graphs/__init__.py +17 -0
  16. contextdb/graphs/base.py +46 -0
  17. contextdb/graphs/causal.py +224 -0
  18. contextdb/graphs/entity.py +251 -0
  19. contextdb/graphs/semantic.py +156 -0
  20. contextdb/graphs/temporal.py +173 -0
  21. contextdb/integrations/__init__.py +10 -0
  22. contextdb/integrations/autogen.py +39 -0
  23. contextdb/integrations/crewai.py +41 -0
  24. contextdb/integrations/langchain.py +132 -0
  25. contextdb/integrations/openai_tools.py +124 -0
  26. contextdb/memory/__init__.py +9 -0
  27. contextdb/memory/experiential.py +102 -0
  28. contextdb/memory/factual.py +58 -0
  29. contextdb/memory/working.py +90 -0
  30. contextdb/privacy/__init__.py +9 -0
  31. contextdb/privacy/audit.py +199 -0
  32. contextdb/privacy/pii_detector.py +173 -0
  33. contextdb/privacy/retention.py +99 -0
  34. contextdb/py.typed +0 -0
  35. contextdb/store/__init__.py +15 -0
  36. contextdb/store/base.py +67 -0
  37. contextdb/store/sqlite_store.py +517 -0
  38. contextdb/store/vector_index.py +241 -0
  39. contextdb/utils/__init__.py +22 -0
  40. contextdb/utils/embeddings.py +159 -0
  41. contextdb/utils/llm.py +139 -0
  42. contextdb/utils/migrations.py +159 -0
  43. pycontextdb-0.1.0.dist-info/METADATA +589 -0
  44. pycontextdb-0.1.0.dist-info/RECORD +47 -0
  45. pycontextdb-0.1.0.dist-info/WHEEL +4 -0
  46. pycontextdb-0.1.0.dist-info/entry_points.txt +2 -0
  47. pycontextdb-0.1.0.dist-info/licenses/LICENSE +190 -0
@@ -0,0 +1,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
+ )