Skip to content

Commit 95f444f

Browse files
committed
perf: set up providers
1 parent 003a820 commit 95f444f

2 files changed

Lines changed: 50 additions & 1 deletion

File tree

backend/open_webui/jms/chat.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,9 +12,14 @@
1212
logger.setLevel(SRC_LOG_LEVELS["WISP"])
1313

1414
CHAT_URL = "/api/v1/terminal/chats/"
15+
TERMINAL_CONFIG_URL = "/api/v1/terminal/terminals/config/"
1516

1617

1718
class ChatHandler(BaseWisp):
19+
def get_providers(self):
20+
resp = self._request("GET", TERMINAL_CONFIG_URL, action="list providers")
21+
return self._loads(TERMINAL_CONFIG_URL, resp.body).get("CHAT_AI_PROVIDERS", [])
22+
1823
def list(self, query: Optional[Dict[str, Any]] = None) -> List[dict]:
1924
query = self._normalize_query(query)
2025
resp = self._request("GET", CHAT_URL, query=query, action="list chats")

backend/open_webui/main.py

Lines changed: 45 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import mimetypes
44
import os
55
import sys
6+
from collections import defaultdict
67

78
from contextlib import asynccontextmanager
89
import anyio.to_thread
@@ -19,7 +20,7 @@
1920
from starlette.exceptions import HTTPException as StarletteHTTPException
2021
from starlette.datastructures import Headers
2122

22-
from open_webui.jms import setup_poll_jms_event
23+
from open_webui.jms import setup_poll_jms_event, chat_manager
2324
from open_webui.utils.logger import start_logger
2425
from open_webui.socket.main import (
2526
app as socket_app,
@@ -108,6 +109,46 @@ async def get_response(self, path: str, scope):
108109
)
109110

110111

112+
def apply_provider_config(providers, config):
113+
grouped = defaultdict(list)
114+
for p in providers:
115+
t = p.get("type")
116+
if t:
117+
grouped[t].append(p)
118+
119+
config_map = {
120+
"openai": (
121+
"OPENAI_API_BASE_URLS",
122+
"OPENAI_API_KEYS",
123+
"OPENAI_API_PROXYS",
124+
),
125+
"ollama": (
126+
"OLLAMA_BASE_URLS",
127+
"OLLAMA_API_KEYS",
128+
"OLLAMA_API_PROXYS",
129+
),
130+
}
131+
132+
for provider_type, (base_attr, key_attr, proxy_attr) in config_map.items():
133+
base_urls = []
134+
api_keys = []
135+
proxys = []
136+
137+
for p in grouped.get(provider_type, []):
138+
base_url = p.get("base_url")
139+
api_key = p.get("api_key")
140+
if not base_url or not api_key:
141+
continue
142+
base_urls.append(base_url)
143+
api_keys.append(api_key)
144+
proxys.append(p.get("proxy"))
145+
146+
if base_urls:
147+
setattr(config, base_attr, base_urls)
148+
setattr(config, key_attr, api_keys)
149+
setattr(config, proxy_attr, proxys)
150+
151+
111152
@asynccontextmanager
112153
async def lifespan(app: FastAPI):
113154
app.state.instance_id = INSTANCE_ID
@@ -166,6 +207,9 @@ async def lifespan(app: FastAPI):
166207
)
167208

168209
setup_poll_jms_event()
210+
211+
providers = chat_manager.get_providers()
212+
apply_provider_config(providers, app.state.config)
169213
yield
170214

171215
if hasattr(app.state, "redis_task_command_listener"):

0 commit comments

Comments
 (0)