@@ -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
0 commit comments