|
3 | 3 | import mimetypes |
4 | 4 | import os |
5 | 5 | import sys |
| 6 | +from collections import defaultdict |
6 | 7 |
|
7 | 8 | from contextlib import asynccontextmanager |
8 | 9 | import anyio.to_thread |
|
19 | 20 | from starlette.exceptions import HTTPException as StarletteHTTPException |
20 | 21 | from starlette.datastructures import Headers |
21 | 22 |
|
22 | | -from open_webui.jms import setup_poll_jms_event |
| 23 | +from open_webui.jms import setup_poll_jms_event, chat_manager |
23 | 24 | from open_webui.utils.logger import start_logger |
24 | 25 | from open_webui.socket.main import ( |
25 | 26 | app as socket_app, |
@@ -108,6 +109,46 @@ async def get_response(self, path: str, scope): |
108 | 109 | ) |
109 | 110 |
|
110 | 111 |
|
| 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 | + |
111 | 152 | @asynccontextmanager |
112 | 153 | async def lifespan(app: FastAPI): |
113 | 154 | app.state.instance_id = INSTANCE_ID |
@@ -166,6 +207,9 @@ async def lifespan(app: FastAPI): |
166 | 207 | ) |
167 | 208 |
|
168 | 209 | setup_poll_jms_event() |
| 210 | + |
| 211 | + providers = chat_manager.get_providers() |
| 212 | + apply_provider_config(providers, app.state.config) |
169 | 213 | yield |
170 | 214 |
|
171 | 215 | if hasattr(app.state, "redis_task_command_listener"): |
|
0 commit comments