Skip to content

Commit 7bae9de

Browse files
committed
Add OpenAI model setup and configuration updates
- Implemented methods to retrieve model definitions and infer embedding dimensions for OpenAI-compatible models. - Enhanced the `_upsert_openai_models` method to update or create model entries in the configuration, ensuring defaults are set correctly. - Updated the setup process to prompt users for custom API settings, including base URL and model names, improving flexibility in configuration. - Added unit tests to validate the OpenAI setup flow, ensuring proper handling of custom API initialization and configuration updates.
1 parent 8c50220 commit 7bae9de

2 files changed

Lines changed: 246 additions & 4 deletions

File tree

backends/advanced/init.py

Lines changed: 110 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -164,6 +164,92 @@ def prompt_with_existing_masked(self, prompt_text: str, env_key: str, placeholde
164164
default=default
165165
)
166166

167+
def _get_model_def(self, config: Dict[str, Any], model_name: str) -> Dict[str, Any]:
168+
"""Get a model definition by name from config.yml."""
169+
models = config.get("models", [])
170+
if not isinstance(models, list):
171+
return {}
172+
return next((m for m in models if m.get("name") == model_name), {})
173+
174+
def _infer_embedding_dimensions(self, model_name: str, fallback: int = 1536) -> int:
175+
"""Infer embedding dimensions for common models."""
176+
known_dimensions = {
177+
"text-embedding-3-small": 1536,
178+
"text-embedding-3-large": 3072,
179+
"text-embedding-ada-002": 1536,
180+
"nomic-embed-text-v1.5": 768,
181+
"nomic-embed-text:latest": 768,
182+
}
183+
return known_dimensions.get(model_name, fallback)
184+
185+
def _upsert_openai_models(
186+
self,
187+
api_key: str,
188+
base_url: str,
189+
llm_model_name: str,
190+
embedding_model_name: str,
191+
) -> None:
192+
"""Update or create openai-llm/openai-embed in config.yml and set defaults."""
193+
config = self.config_manager.get_full_config()
194+
models = config.get("models", [])
195+
if not isinstance(models, list):
196+
models = []
197+
198+
openai_llm = self._get_model_def(config, "openai-llm")
199+
openai_embed = self._get_model_def(config, "openai-embed")
200+
201+
llm_params = openai_llm.get("model_params", {})
202+
if not isinstance(llm_params, dict):
203+
llm_params = {}
204+
llm_params.setdefault("temperature", 0.2)
205+
llm_params.setdefault("max_tokens", 2000)
206+
207+
embedding_dimensions = openai_embed.get("embedding_dimensions")
208+
if not isinstance(embedding_dimensions, int) or embedding_dimensions <= 0:
209+
embedding_dimensions = self._infer_embedding_dimensions(embedding_model_name)
210+
211+
llm_payload = {
212+
"name": "openai-llm",
213+
"description": "OpenAI/OpenAI-compatible LLM",
214+
"model_type": "llm",
215+
"model_provider": "openai",
216+
"api_family": "openai",
217+
"model_name": llm_model_name,
218+
"model_url": base_url,
219+
"api_key": api_key,
220+
"model_params": llm_params,
221+
"model_output": "json",
222+
}
223+
embed_payload = {
224+
"name": "openai-embed",
225+
"description": "OpenAI/OpenAI-compatible embeddings",
226+
"model_type": "embedding",
227+
"model_provider": "openai",
228+
"api_family": "openai",
229+
"model_name": embedding_model_name,
230+
"model_url": base_url,
231+
"api_key": api_key,
232+
"embedding_dimensions": embedding_dimensions,
233+
"model_output": "vector",
234+
}
235+
236+
def upsert_model(payload: Dict[str, Any]):
237+
for idx, model in enumerate(models):
238+
if model.get("name") == payload["name"]:
239+
models[idx] = {**model, **payload}
240+
return
241+
models.append(payload)
242+
243+
upsert_model(llm_payload)
244+
upsert_model(embed_payload)
245+
246+
config["models"] = models
247+
if "defaults" not in config or not isinstance(config["defaults"], dict):
248+
config["defaults"] = {}
249+
config["defaults"]["llm"] = "openai-llm"
250+
config["defaults"]["embedding"] = "openai-embed"
251+
252+
self.config_manager.save_full_config(config)
167253

168254
def setup_authentication(self):
169255
"""Configure authentication settings"""
@@ -307,17 +393,29 @@ def setup_llm(self):
307393
self.console.print()
308394

309395
choices = {
310-
"1": "OpenAI (GPT-4, GPT-3.5 - requires API key)",
396+
"1": "OpenAI / OpenAI-compatible (custom base URL, API key, model names)",
311397
"2": "Ollama (local models - runs locally)",
312398
"3": "Skip (no memory extraction)"
313399
}
314400

315401
choice = self.prompt_choice("Which LLM provider will you use?", choices, "1")
316402

317403
if choice == "1":
318-
self.console.print("[blue][INFO][/blue] OpenAI selected")
404+
self.console.print("[blue][INFO][/blue] OpenAI/OpenAI-compatible selected")
319405
self.console.print("Get your API key from: https://platform.openai.com/api-keys")
320406

407+
existing_cfg = self.config_manager.get_full_config()
408+
openai_llm = self._get_model_def(existing_cfg, "openai-llm")
409+
openai_embed = self._get_model_def(existing_cfg, "openai-embed")
410+
411+
default_base_url = openai_llm.get("model_url") or openai_embed.get("model_url") or "https://api.openai.com/v1"
412+
default_llm_model = openai_llm.get("model_name") or "gpt-4o-mini"
413+
default_embedding_model = openai_embed.get("model_name") or "text-embedding-3-small"
414+
415+
base_url = self.prompt_value("OpenAI-compatible base URL", default_base_url)
416+
llm_model_name = self.prompt_value("LLM model name", default_llm_model)
417+
embedding_model_name = self.prompt_value("Embedding model name", default_embedding_model)
418+
321419
# Use the new masked prompt function
322420
api_key = self.prompt_with_existing_masked(
323421
prompt_text="OpenAI API key (leave empty to skip)",
@@ -329,11 +427,19 @@ def setup_llm(self):
329427

330428
if api_key:
331429
self.config["OPENAI_API_KEY"] = api_key
332-
# Update config.yml to use OpenAI models
333-
self.config_manager.update_config_defaults({"llm": "openai-llm", "embedding": "openai-embed"})
430+
# Update config.yml openai model definitions and defaults
431+
self._upsert_openai_models(
432+
api_key=api_key,
433+
base_url=base_url,
434+
llm_model_name=llm_model_name,
435+
embedding_model_name=embedding_model_name,
436+
)
334437
self.console.print("[green][SUCCESS][/green] OpenAI configured in config.yml")
335438
self.console.print("[blue][INFO][/blue] Set defaults.llm: openai-llm")
336439
self.console.print("[blue][INFO][/blue] Set defaults.embedding: openai-embed")
440+
self.console.print(f"[blue][INFO][/blue] Set openai-llm.model_url: {base_url}")
441+
self.console.print(f"[blue][INFO][/blue] Set openai-llm.model_name: {llm_model_name}")
442+
self.console.print(f"[blue][INFO][/blue] Set openai-embed.model_name: {embedding_model_name}")
337443
else:
338444
self.console.print("[yellow][WARNING][/yellow] No API key provided - memory extraction will not work")
339445

Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,136 @@
1+
"""Unit tests for OpenAI custom API setup/initialization flow in init.py.
2+
3+
These tests verify that wizard setup can initialize OpenAI-compatible providers
4+
with custom API endpoints (base URL), API keys, and model names.
5+
"""
6+
7+
import importlib.util
8+
from pathlib import Path
9+
from unittest.mock import MagicMock, call
10+
11+
12+
_INIT_PATH = Path(__file__).resolve().parents[1] / "init.py"
13+
_SPEC = importlib.util.spec_from_file_location("advanced_init", _INIT_PATH)
14+
_MODULE = importlib.util.module_from_spec(_SPEC)
15+
assert _SPEC and _SPEC.loader
16+
_SPEC.loader.exec_module(_MODULE)
17+
ChronicleSetup = _MODULE.ChronicleSetup
18+
19+
20+
def _build_setup_with_mocks() -> ChronicleSetup:
21+
"""Create ChronicleSetup instance without running __init__ side effects."""
22+
setup = ChronicleSetup.__new__(ChronicleSetup)
23+
setup.console = MagicMock()
24+
setup.config = {}
25+
setup.config_manager = MagicMock()
26+
setup.print_section = MagicMock()
27+
return setup
28+
29+
30+
def test_upsert_openai_models_updates_existing_defs_and_defaults():
31+
"""Checks OpenAI custom API config upsert in init flow.
32+
33+
Verifies that init setup updates OpenAI model definitions with custom API
34+
settings and switches defaults to those OpenAI entries.
35+
"""
36+
setup = _build_setup_with_mocks()
37+
setup.config_manager.get_full_config.return_value = {
38+
"defaults": {"llm": "local-llm", "embedding": "local-embed"},
39+
"models": [
40+
{
41+
"name": "openai-llm",
42+
"model_type": "llm",
43+
"model_provider": "openai",
44+
"model_name": "gpt-4o-mini",
45+
"model_url": "https://api.openai.com/v1",
46+
"api_key": "old-key",
47+
"model_params": {"temperature": 0.3},
48+
},
49+
{
50+
"name": "openai-embed",
51+
"model_type": "embedding",
52+
"model_provider": "openai",
53+
"model_name": "text-embedding-3-small",
54+
"model_url": "https://api.openai.com/v1",
55+
"api_key": "old-key",
56+
"embedding_dimensions": 1536,
57+
},
58+
],
59+
}
60+
61+
setup._upsert_openai_models(
62+
api_key="new-key",
63+
base_url="http://custom.example/v1",
64+
llm_model_name="gpt-oss-20b",
65+
embedding_model_name="text-embedding-3-large",
66+
)
67+
68+
saved_config = setup.config_manager.save_full_config.call_args[0][0]
69+
saved_models = {m["name"]: m for m in saved_config["models"]}
70+
71+
assert saved_config["defaults"]["llm"] == "openai-llm"
72+
assert saved_config["defaults"]["embedding"] == "openai-embed"
73+
74+
assert saved_models["openai-llm"]["model_url"] == "http://custom.example/v1"
75+
assert saved_models["openai-llm"]["api_key"] == "new-key"
76+
assert saved_models["openai-llm"]["model_name"] == "gpt-oss-20b"
77+
# Existing params are preserved and missing defaults are filled.
78+
assert saved_models["openai-llm"]["model_params"]["temperature"] == 0.3
79+
assert saved_models["openai-llm"]["model_params"]["max_tokens"] == 2000
80+
81+
assert saved_models["openai-embed"]["model_url"] == "http://custom.example/v1"
82+
assert saved_models["openai-embed"]["api_key"] == "new-key"
83+
assert saved_models["openai-embed"]["model_name"] == "text-embedding-3-large"
84+
# Existing embedding dimensions are preserved.
85+
assert saved_models["openai-embed"]["embedding_dimensions"] == 1536
86+
87+
88+
def test_setup_llm_openai_prompts_for_custom_values_and_updates_models():
89+
"""Checks init OpenAI setup prompts for custom API initialization values."""
90+
setup = _build_setup_with_mocks()
91+
setup.prompt_choice = MagicMock(return_value="1")
92+
setup.prompt_value = MagicMock(
93+
side_effect=["http://my-openai-compatible/v1", "my-chat-model", "my-embed-model"]
94+
)
95+
setup.prompt_with_existing_masked = MagicMock(return_value="test-api-key")
96+
setup._upsert_openai_models = MagicMock()
97+
setup.config_manager.get_full_config.return_value = {
98+
"models": [
99+
{"name": "openai-llm", "model_url": "https://api.openai.com/v1", "model_name": "gpt-4o-mini"},
100+
{"name": "openai-embed", "model_url": "https://api.openai.com/v1", "model_name": "text-embedding-3-small"},
101+
]
102+
}
103+
104+
setup.setup_llm()
105+
106+
setup.prompt_value.assert_has_calls(
107+
[
108+
call("OpenAI-compatible base URL", "https://api.openai.com/v1"),
109+
call("LLM model name", "gpt-4o-mini"),
110+
call("Embedding model name", "text-embedding-3-small"),
111+
]
112+
)
113+
setup._upsert_openai_models.assert_called_once_with(
114+
api_key="test-api-key",
115+
base_url="http://my-openai-compatible/v1",
116+
llm_model_name="my-chat-model",
117+
embedding_model_name="my-embed-model",
118+
)
119+
assert setup.config["OPENAI_API_KEY"] == "test-api-key"
120+
121+
122+
def test_setup_llm_openai_skips_upsert_when_api_key_missing():
123+
"""Checks init OpenAI custom API setup guards against missing API key."""
124+
setup = _build_setup_with_mocks()
125+
setup.prompt_choice = MagicMock(return_value="1")
126+
setup.prompt_value = MagicMock(
127+
side_effect=["https://api.openai.com/v1", "gpt-4o-mini", "text-embedding-3-small"]
128+
)
129+
setup.prompt_with_existing_masked = MagicMock(return_value="")
130+
setup._upsert_openai_models = MagicMock()
131+
setup.config_manager.get_full_config.return_value = {"models": []}
132+
133+
setup.setup_llm()
134+
135+
setup._upsert_openai_models.assert_not_called()
136+
assert "OPENAI_API_KEY" not in setup.config

0 commit comments

Comments
 (0)