From 6609918bfeebf9b8b194cfeef9dda6f06128677e Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Mon, 31 Aug 2026 00:53:55 -0400 Subject: [PATCH] refac --- backend/open_webui/routers/ollama.py | 5 +++++ backend/open_webui/routers/openai.py | 14 +++++------- backend/open_webui/utils/models.py | 33 +++++++++++++++++++++------- 3 files changed, 36 insertions(+), 16 deletions(-) diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 2704d81d1b..e9a4685993 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -25,6 +25,7 @@ from open_webui.env import ( ENABLE_FORWARD_USER_INFO_HEADERS, FORWARD_SESSION_INFO_HEADER_CHAT_ID, MODELS_CACHE_TTL, + REDIS_KEY_PREFIX, ) from open_webui.events import EVENTS, publish_event, publish_model_provider_request_failed from open_webui.internal.db import get_async_session @@ -56,6 +57,7 @@ log = logging.getLogger(__name__) # in ZlibError. See https://github.com/aio-libs/aiohttp/issues/4462. _STRIP_PROXY_HEADERS = frozenset({'Content-Encoding', 'Content-Length', 'Transfer-Encoding'}) _MODEL_LIST_TIMEOUT = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST) +BASE_MODELS_CACHE_KEY = f'{REDIS_KEY_PREFIX}:models:base' def _clean_proxy_headers(raw_headers) -> dict: @@ -319,6 +321,9 @@ async def update_config( ) await get_all_models.cache.clear() + redis = getattr(request.app.state, 'redis', None) + if redis is not None: + await redis.delete(BASE_MODELS_CACHE_KEY) request.app.state.BASE_MODELS = [] request.app.state.OLLAMA_MODELS = {} models = getattr(request.app.state, 'MODELS', None) diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 0f21defdd5..d7856bea48 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -30,6 +30,7 @@ from open_webui.env import ( ENABLE_OPENAI_API_PASSTHROUGH, FORWARD_SESSION_INFO_HEADER_CHAT_ID, MODELS_CACHE_TTL, + REDIS_KEY_PREFIX, ) from open_webui.events import EVENTS, publish_event, publish_model_provider_request_failed from open_webui.internal.db import get_async_session @@ -76,6 +77,7 @@ log = logging.getLogger(__name__) _STRIP_PROXY_HEADERS = frozenset({'Content-Encoding', 'Content-Length', 'Transfer-Encoding'}) _MODEL_LIST_TIMEOUT = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST) _UNSUPPORTED_OPENAI_MODEL_KEYWORDS = ('babbage', 'dall-e', 'davinci', 'embedding', 'tts', 'whisper') +BASE_MODELS_CACHE_KEY = f'{REDIS_KEY_PREFIX}:models:base' def _clean_proxy_headers(raw_headers) -> dict: @@ -345,6 +347,9 @@ async def get_openai_connection(idx: int) -> tuple[str, str, dict]: async def clear_openai_model_cache(request: Request): await get_all_models.cache.clear() + redis = getattr(request.app.state, 'redis', None) + if redis is not None: + await redis.delete(BASE_MODELS_CACHE_KEY) request.app.state.BASE_MODELS = [] request.app.state.OPENAI_MODELS = {} models = getattr(request.app.state, 'MODELS', None) @@ -571,14 +576,7 @@ async def update_config(request: Request, form_data: OpenAIConfigForm, user=Depe } ) - await get_all_models.cache.clear() - request.app.state.BASE_MODELS = [] - request.app.state.OPENAI_MODELS = {} - models = getattr(request.app.state, 'MODELS', None) - if hasattr(models, 'clear'): - models.clear() - else: - request.app.state.MODELS = {} + await clear_openai_model_cache(request) await publish_event( request, diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index 717718cc99..4e1559200a 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -3,13 +3,12 @@ import copy import logging import sys -from aiocache import cached from fastapi import Request from open_webui.config import ( BYPASS_ADMIN_ACCESS_CONTROL, DEFAULT_ARENA_MODEL, ) -from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_PLUGINS, GLOBAL_LOG_LEVEL +from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_PLUGINS, GLOBAL_LOG_LEVEL, REDIS_KEY_PREFIX from open_webui.functions import get_function_models from open_webui.models.access_grants import AccessGrants from open_webui.models.config import Config @@ -21,6 +20,7 @@ from open_webui.models.users import UserModel from open_webui.routers import ollama, openai from open_webui.socket.utils import RedisDict from open_webui.utils.access_control import has_access, has_base_model_access +from open_webui.utils.json_codec import JSONCodec from open_webui.utils.plugin import ( get_functions_cache, get_function_module_from_cache, @@ -29,6 +29,8 @@ from open_webui.utils.plugin import ( logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) log = logging.getLogger(__name__) +BASE_MODELS_CACHE_KEY = f'{REDIS_KEY_PREFIX}:models:base' + async def fetch_ollama_models(request: Request, user: UserModel = None): raw_ollama_models = await ollama.get_all_models(request, user=user) @@ -74,17 +76,32 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) if refresh: await openai.get_all_models.cache.clear() await ollama.get_all_models.cache.clear() + redis = getattr(request.app.state, 'redis', None) + if redis is not None: + await redis.delete(BASE_MODELS_CACHE_KEY) + request.app.state.BASE_MODELS = [] - if ( - request.app.state.MODELS - and request.app.state.BASE_MODELS - and (config.get('models.base_models_cache') and not refresh) - ): + redis = getattr(request.app.state, 'redis', None) + use_cache = config.get('models.base_models_cache') and not refresh + base_models = None + + if use_cache and redis is not None: + cached_base_models = await redis.get(BASE_MODELS_CACHE_KEY) + if cached_base_models: + base_models = JSONCodec.loads(cached_base_models) + request.app.state.BASE_MODELS = base_models + else: + await openai.get_all_models.cache.clear() + await ollama.get_all_models.cache.clear() + elif use_cache and request.app.state.MODELS and request.app.state.BASE_MODELS: base_models = request.app.state.BASE_MODELS - else: + + if base_models is None: base_models = await get_all_base_models(request, user=user) if base_models: request.app.state.BASE_MODELS = base_models + if config.get('models.base_models_cache') and redis is not None: + await redis.set(BASE_MODELS_CACHE_KEY, JSONCodec.dumps(base_models)) else: base_models = request.app.state.BASE_MODELS