mirror of
https://github.com/open-webui/open-webui.git
synced 2026-09-02 12:14:48 +02:00
refac
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user