local models docker url

frdel committed Dec 5, 2024 at 20:45 UTC 21975c5a7cc7b3ad8b9ab95f940b5e6f6a743231
2 files changed +39 -12
models.py
+34 -12
@@ -26,7 +26,7 @@ from langchain_google_genai import (
26 )
27 from langchain_mistralai import ChatMistralAI
28 from pydantic.v1.types import SecretStr
29 -from python.helpers import dotenv
29 +from python.helpers import dotenv, runtime
30 from python.helpers.dotenv import load_dotenv
31
32 # environment variables
@@ -71,7 +71,12 @@ def get_model(type: ModelType, provider: ModelProvider, name: str, **kwargs):
71 return model
72
73
74 +
75 +
76 # Ollama models
77 +def get_ollama_base_url():
78 + return dotenv.get_dotenv_value("OLLAMA_BASE_URL") or f"http://{runtime.get_local_url()}:11434"
79 +
80 def get_ollama_chat(
81 model_name: str,
82 temperature=DEFAULT_TEMPERATURE,
@@ -80,7 +85,7 @@ def get_ollama_chat(
85 **kwargs,
86 ):
87 if not base_url:
83 - base_url = dotenv.get_dotenv_value("OLLAMA_BASE_URL") or "http://127.0.0.1:11434"
88 + base_url = get_ollama_base_url()
89 return ChatOllama(
90 model=model_name,
91 temperature=temperature,
@@ -97,7 +102,7 @@ def get_ollama_embedding(
102 **kwargs,
103 ):
104 if not base_url:
100 - base_url = dotenv.get_dotenv_value("OLLAMA_BASE_URL") or "http://127.0.0.1:11434"
105 + base_url = get_ollama_base_url()
106 return OllamaEmbeddings(
107 model=model_name, temperature=temperature, base_url=base_url, **kwargs
108 )
@@ -132,22 +137,27 @@ def get_huggingface_embedding(model_name: str, **kwargs):
137
138
139 # LM Studio and other OpenAI compatible interfaces
140 +def get_lmstudio_base_url():
141 + return dotenv.get_dotenv_value("LM_STUDIO_BASE_URL") or f"http://{runtime.get_local_url()}:1234/v1"
142 +
143 def get_lmstudio_chat(
144 model_name: str,
145 temperature=DEFAULT_TEMPERATURE,
138 - base_url=dotenv.get_dotenv_value("LM_STUDIO_BASE_URL")
139 - or "http://127.0.0.1:1234/v1",
146 + base_url=None,
147 **kwargs,
148 ):
149 + if not base_url:
150 + base_url = get_lmstudio_base_url()
151 return ChatOpenAI(model_name=model_name, base_url=base_url, temperature=temperature, api_key="none", **kwargs) # type: ignore
152
153
154 def get_lmstudio_embedding(
155 model_name: str,
147 - base_url=dotenv.get_dotenv_value("LM_STUDIO_BASE_URL")
148 - or "http://127.0.0.1:1234/v1",
156 + base_url=None,
157 **kwargs,
158 ):
159 + if not base_url:
160 + base_url = get_lmstudio_base_url()
161 return OpenAIEmbeddings(model=model_name, api_key="none", base_url=base_url, check_embedding_ctx_length=False, **kwargs) # type: ignore
162
163
@@ -227,7 +237,7 @@ def get_google_chat(
237 **kwargs,
238 ):
239 if not api_key:
230 - api_key = get_api_key("google")
240 + api_key = get_api_key("google")
241 return GoogleGenerativeAI(model=model_name, temperature=temperature, google_api_key=api_key, safety_settings={HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT: HarmBlockThreshold.BLOCK_NONE}, **kwargs) # type: ignore
242
243
@@ -276,7 +286,10 @@ def get_openrouter_chat(
286 if not api_key:
287 api_key = get_api_key("openrouter")
288 if not base_url:
279 - base_url = dotenv.get_dotenv_value("OPEN_ROUTER_BASE_URL") or "https://openrouter.ai/api/v1"
289 + base_url = (
290 + dotenv.get_dotenv_value("OPEN_ROUTER_BASE_URL")
291 + or "https://openrouter.ai/api/v1"
292 + )
293 return ChatOpenAI(api_key=api_key, model=model_name, temperature=temperature, base_url=base_url, **kwargs) # type: ignore
294
295
@@ -289,7 +302,10 @@ def get_openrouter_embedding(
302 if not api_key:
303 api_key = get_api_key("openrouter")
304 if not base_url:
292 - base_url = dotenv.get_dotenv_value("OPEN_ROUTER_BASE_URL") or "https://openrouter.ai/api/v1"
305 + base_url = (
306 + dotenv.get_dotenv_value("OPEN_ROUTER_BASE_URL")
307 + or "https://openrouter.ai/api/v1"
308 + )
309 return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url, **kwargs) # type: ignore
310
311
@@ -305,7 +321,10 @@ def get_sambanova_chat(
321 if not api_key:
322 api_key = get_api_key("sambanova")
323 if not base_url:
308 - base_url = dotenv.get_dotenv_value("SAMBANOVA_BASE_URL") or "https://fast-api.snova.ai/v1"
324 + base_url = (
325 + dotenv.get_dotenv_value("SAMBANOVA_BASE_URL")
326 + or "https://fast-api.snova.ai/v1"
327 + )
328 return ChatOpenAI(api_key=api_key, model=model_name, temperature=temperature, base_url=base_url, max_tokens=max_tokens, **kwargs) # type: ignore
329
330
@@ -319,7 +338,10 @@ def get_sambanova_embedding(
338 if not api_key:
339 api_key = get_api_key("sambanova")
340 if not base_url:
322 - base_url = dotenv.get_dotenv_value("SAMBANOVA_BASE_URL") or "https://fast-api.snova.ai/v1"
341 + base_url = (
342 + dotenv.get_dotenv_value("SAMBANOVA_BASE_URL")
343 + or "https://fast-api.snova.ai/v1"
344 + )
345 return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url, **kwargs) # type: ignore
346
347
python/helpers/runtime.py
+5
@@ -47,6 +47,11 @@ def is_dockerized() -> bool:
47 def is_development() -> bool:
48 return not is_dockerized()
49
50 +def get_local_url():
51 + if is_dockerized():
52 + return "host.docker.internal"
53 + return "127.0.0.1"
54 +
55 async def call_development_function(func: Callable, *args, **kwargs):
56 if is_development():
57 url = _get_rfc_url()