Coverage for open_webui/utils/models.py: 16%
308 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import asyncio
2import copy
3import logging
4import sys
6from fastapi import Request
7from open_webui.config import (
8 BYPASS_ADMIN_ACCESS_CONTROL,
9 DEFAULT_ARENA_MODEL,
10)
11from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_PLUGINS, GLOBAL_LOG_LEVEL, REDIS_KEY_PREFIX
12from open_webui.functions import get_function_models
13from open_webui.models.access_grants import AccessGrants
14from open_webui.models.config import Config
15from open_webui.models.functions import Functions
16from open_webui.models.groups import Groups
17from open_webui.models.models import Models
18from open_webui.utils.chat_variables import get_chat_variables_schema
19from open_webui.models.users import UserModel
20from open_webui.routers import ollama, openai
21from open_webui.socket.utils import RedisDict
22from open_webui.utils.access_control import has_access, has_base_model_access
23from open_webui.utils.json_codec import JSONCodec
24from open_webui.utils.plugin import (
25 get_functions_cache,
26 get_function_module_from_cache,
27)
29logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
30log = logging.getLogger(__name__)
32BASE_MODELS_CACHE_KEY = f'{REDIS_KEY_PREFIX}:models:base'
35async def fetch_ollama_models(request: Request, user: UserModel = None):
36 raw_ollama_models = await ollama.get_all_models(request, user=user)
37 return [
38 {
39 'id': model['model'],
40 'name': model['name'],
41 'object': 'model',
42 'created': 0,
43 'owned_by': 'ollama',
44 'ollama': model,
45 'loaded': 'expires_at' in model,
46 'connection_type': model.get('connection_type', 'local'),
47 'tags': model.get('tags', []),
48 }
49 for model in raw_ollama_models['models']
50 ]
53async def fetch_openai_models(request: Request, user: UserModel = None):
54 openai_response = await openai.get_all_models(request, user=user)
55 return openai_response['data']
58async def get_all_base_models(request: Request, user: UserModel = None):
59 config = await Config.get_many('openai.enable', 'ollama.enable')
60 openai_task = fetch_openai_models(request, user) if config.get('openai.enable') else asyncio.sleep(0, result=[])
61 ollama_task = fetch_ollama_models(request, user) if config.get('ollama.enable') else asyncio.sleep(0, result=[])
62 function_task = get_function_models(request)
64 openai_models, ollama_models, function_models = await asyncio.gather(openai_task, ollama_task, function_task)
66 return function_models + openai_models + ollama_models
69async def get_all_models(request, refresh: bool = False, user: UserModel = None):
70 config = await Config.get_many(
71 'models.base_models_cache',
72 'evaluation.arena.enable',
73 'evaluation.arena.models',
74 'models.default_metadata',
75 )
76 if refresh:
77 await openai.get_all_models.cache.clear()
78 await ollama.get_all_models.cache.clear()
79 redis = getattr(request.app.state, 'redis', None)
80 if redis is not None: 80 ↛ 81line 80 didn't jump to line 81 because the condition on line 80 was never true
81 await redis.delete(BASE_MODELS_CACHE_KEY)
82 request.app.state.BASE_MODELS = []
84 redis = getattr(request.app.state, 'redis', None)
85 use_cache = config.get('models.base_models_cache') and not refresh
86 base_models = None
88 if use_cache and redis is not None: 88 ↛ 89line 88 didn't jump to line 89 because the condition on line 88 was never true
89 cached_base_models = await redis.get(BASE_MODELS_CACHE_KEY)
90 if cached_base_models:
91 base_models = JSONCodec.loads(cached_base_models)
92 request.app.state.BASE_MODELS = base_models
93 else:
94 await openai.get_all_models.cache.clear()
95 await ollama.get_all_models.cache.clear()
96 elif use_cache and request.app.state.MODELS and request.app.state.BASE_MODELS: 96 ↛ 97line 96 didn't jump to line 97 because the condition on line 96 was never true
97 base_models = request.app.state.BASE_MODELS
99 if base_models is None: 99 ↛ 109line 99 didn't jump to line 109 because the condition on line 99 was always true
100 base_models = await get_all_base_models(request, user=user)
101 if base_models: 101 ↛ 102line 101 didn't jump to line 102 because the condition on line 101 was never true
102 request.app.state.BASE_MODELS = base_models
103 if config.get('models.base_models_cache') and redis is not None:
104 await redis.set(BASE_MODELS_CACHE_KEY, JSONCodec.dumps(base_models))
105 else:
106 base_models = request.app.state.BASE_MODELS
108 # deep copy the base models to avoid modifying the original list
109 models = [model.copy() for model in base_models]
111 # If there are no models, return an empty list
112 if len(models) == 0: 112 ↛ 116line 112 didn't jump to line 116 because the condition on line 112 was always true
113 return []
115 # Add arena models
116 if config.get('evaluation.arena.enable'):
117 arena_models = []
118 arena_config = config.get('evaluation.arena.models') or []
119 if len(arena_config) > 0:
120 arena_models = [
121 {
122 'id': model['id'],
123 'name': model['name'],
124 'info': {
125 'meta': model['meta'],
126 },
127 'object': 'model',
128 'created': 0,
129 'owned_by': 'arena',
130 'arena': True,
131 }
132 for model in arena_config
133 ]
134 else:
135 # Add default arena model
136 arena_models = [
137 {
138 'id': DEFAULT_ARENA_MODEL['id'],
139 'name': DEFAULT_ARENA_MODEL['name'],
140 'info': {
141 'meta': DEFAULT_ARENA_MODEL['meta'],
142 },
143 'object': 'model',
144 'created': 0,
145 'owned_by': 'arena',
146 'arena': True,
147 }
148 ]
149 models = models + arena_models
151 # One query per type: the global sets are subsets of the active sets, so
152 # deriving them from the same rows halves the function-table queries.
153 if ENABLE_PLUGINS:
154 active_actions = await Functions.get_active_function_ids_by_type('action')
155 global_action_ids = {function_id for function_id, is_global in active_actions if is_global}
156 enabled_action_ids = {function_id for function_id, _ in active_actions}
158 active_filters = await Functions.get_active_function_ids_by_type('filter')
159 global_filter_ids = {function_id for function_id, is_global in active_filters if is_global}
160 enabled_filter_ids = {function_id for function_id, _ in active_filters}
161 else:
162 global_action_ids = set()
163 enabled_action_ids = set()
164 global_filter_ids = set()
165 enabled_filter_ids = set()
167 custom_models = await Models.get_all_models()
169 # Single O(1) lookup: Ollama base names first, then exact IDs (exact wins).
170 base_model_lookup = {}
171 for model in models:
172 if model.get('owned_by') == 'ollama':
173 base_model_lookup.setdefault(model['id'].split(':')[0], model)
174 base_model_lookup[model['id']] = model
176 existing_ids = {m['id'] for m in models}
178 for custom_model in custom_models:
179 if custom_model.base_model_id is None:
180 # Override applied directly to a base model (shares the same ID)
181 model = base_model_lookup.get(custom_model.id)
183 if model:
184 if custom_model.is_active:
185 model['name'] = custom_model.name
186 model['info'] = custom_model.model_dump()
187 schema = get_chat_variables_schema(custom_model.params.model_dump().get('system'))
188 if schema:
189 model['info'].setdefault('meta', {})['chat_variables_schema'] = schema
190 elif isinstance(model['info'].get('meta'), dict):
191 model['info']['meta'].pop('chat_variables_schema', None)
193 action_ids = []
194 filter_ids = []
196 if 'info' in model:
197 if 'meta' in model['info']:
198 if ENABLE_PLUGINS:
199 action_ids.extend(model['info']['meta'].get('actionIds', []))
200 filter_ids.extend(model['info']['meta'].get('filterIds', []))
202 if 'params' in model['info']:
203 del model['info']['params']
205 model['action_ids'] = action_ids
206 model['filter_ids'] = filter_ids
207 else:
208 models = [m for m in models if m is not model]
210 elif custom_model.is_active:
211 if custom_model.id in existing_ids:
212 continue
214 owned_by = 'openai'
215 connection_type = None
216 pipe = None
218 base_model = base_model_lookup.get(custom_model.base_model_id)
219 if base_model is None:
220 base_model = base_model_lookup.get(custom_model.base_model_id.split(':')[0])
221 if base_model:
222 owned_by = base_model.get('owned_by', 'unknown')
223 if 'pipe' in base_model:
224 pipe = base_model['pipe']
225 connection_type = base_model.get('connection_type', None)
227 model = {
228 'id': f'{custom_model.id}',
229 'name': custom_model.name,
230 'object': 'model',
231 'created': custom_model.created_at,
232 'owned_by': owned_by,
233 'connection_type': connection_type,
234 'preset': True,
235 **({'pipe': pipe} if pipe is not None else {}),
236 **({'provider': base_model.get('provider')} if base_model and base_model.get('provider') else {}),
237 **({'loaded': base_model.get('loaded')} if base_model and base_model.get('loaded') is not None else {}),
238 }
240 info = custom_model.model_dump()
241 schema = get_chat_variables_schema(custom_model.params.model_dump().get('system'))
242 if schema:
243 info.setdefault('meta', {})['chat_variables_schema'] = schema
244 elif isinstance(info.get('meta'), dict):
245 info['meta'].pop('chat_variables_schema', None)
246 if 'params' in info:
247 # Remove params to avoid exposing sensitive info
248 del info['params']
250 model['info'] = info
252 action_ids = []
253 filter_ids = []
255 if custom_model.meta:
256 meta = custom_model.meta.model_dump()
258 if ENABLE_PLUGINS and 'actionIds' in meta:
259 action_ids.extend(meta['actionIds'])
261 if ENABLE_PLUGINS and 'filterIds' in meta:
262 filter_ids.extend(meta['filterIds'])
264 model['action_ids'] = action_ids
265 model['filter_ids'] = filter_ids
267 models.append(model)
269 # Process action_ids to get the actions
270 def get_action_items_from_module(function, module):
271 actions = []
272 if hasattr(module, 'actions'):
273 actions = module.actions
274 return [
275 {
276 'id': f'{function.id}.{action["id"]}',
277 'name': action.get('name', f'{function.name} ({action["id"]})'),
278 'description': function.meta.description,
279 'icon': action.get(
280 'icon_url',
281 function.meta.manifest.get('icon_url', None)
282 or getattr(module, 'icon_url', None)
283 or getattr(module, 'icon', None),
284 ),
285 }
286 for action in actions
287 ]
288 else:
289 return [
290 {
291 'id': function.id,
292 'name': function.name,
293 'description': function.meta.description,
294 'icon': function.meta.manifest.get('icon_url', None)
295 or getattr(module, 'icon_url', None)
296 or getattr(module, 'icon', None),
297 }
298 ]
300 # Process filter_ids to get the filters
301 def get_filter_items_from_module(function, module):
302 return [
303 {
304 'id': function.id,
305 'name': function.name,
306 'description': function.meta.description,
307 'icon': function.meta.manifest.get('icon_url', None)
308 or getattr(module, 'icon_url', None)
309 or getattr(module, 'icon', None),
310 'has_user_valves': hasattr(module, 'UserValves'),
311 }
312 ]
314 # Batch-prefetch all needed function records to avoid N+1 queries
315 all_function_ids = set()
316 for model in models:
317 all_function_ids.update(model.get('action_ids', []))
318 all_function_ids.update(model.get('filter_ids', []))
319 all_function_ids.update(global_action_ids)
320 all_function_ids.update(global_filter_ids)
322 functions_by_id = {f.id: f for f in await Functions.get_functions_by_ids(list(all_function_ids))}
324 # Pre-warm the function module cache once per unique function ID.
325 # This ensures each function's DB freshness check runs exactly once,
326 # not once per (model × function) pair.
327 # Only attempt to load functions that actually exist in the local DB;
328 # imported/custom model configs may reference tools or filters the user
329 # hasn't installed, and trying to load those would cause persistent
330 # "Failed to load function module" log spam on every model refresh.
331 for function_id, function in functions_by_id.items():
332 try:
333 await get_function_module_from_cache(request, function_id, function=function)
334 except Exception as e:
335 log.debug('Failed to load function module for %s: %s', function_id, e)
337 # Apply global model defaults to all models
338 # Per-model overrides take precedence over global defaults
339 default_metadata = config.get('models.default_metadata') or {}
341 if default_metadata:
342 for model in models:
343 info = model.get('info')
345 if info is None:
346 model['info'] = {'meta': copy.deepcopy(default_metadata)}
347 continue
349 meta = info.setdefault('meta', {})
350 for key, value in default_metadata.items():
351 if key == 'capabilities':
352 # Merge capabilities: defaults as base, per-model overrides win
353 existing = meta.get('capabilities') or {}
354 meta['capabilities'] = {**value, **existing}
355 elif meta.get(key) is None:
356 meta[key] = copy.deepcopy(value)
358 # Batch-fetch all function valves in one query to avoid N+1 DB hits
359 # inside get_action_priority (previously called per action × per model).
360 all_function_valves = await Functions.get_function_valves_by_ids(list(all_function_ids))
361 functions_cache = get_functions_cache(request)
363 # Global actions and filters appear in every model, so priorities and item
364 # lists are memoized across the loop instead of rebuilt per model.
365 action_priorities = {}
367 def get_action_priority(action_id):
368 if action_id in action_priorities:
369 return action_priorities[action_id]
370 priority = 0
371 try:
372 function_module = functions_cache.get(action_id)
373 if function_module and hasattr(function_module, 'Valves'):
374 valves_db = all_function_valves.get(action_id)
375 valves = function_module.Valves(**(valves_db if valves_db else {}))
376 priority = getattr(valves, 'priority', 0)
377 except Exception:
378 priority = 0
379 action_priorities[action_id] = priority
380 return priority
382 action_items_by_id = {}
383 filter_items_by_id = {}
385 for model in models:
386 action_ids = [
387 action_id
388 for action_id in set(model.pop('action_ids', [])) | global_action_ids
389 if action_id in enabled_action_ids
390 ]
391 action_ids.sort(key=lambda aid: (get_action_priority(aid), aid))
393 filter_ids = [
394 filter_id
395 for filter_id in set(model.pop('filter_ids', [])) | global_filter_ids
396 if filter_id in enabled_filter_ids
397 ]
398 # Set order varies per process, and an unstable order defeats the RedisDict content signature.
399 filter_ids.sort()
401 model['actions'] = []
402 for action_id in action_ids:
403 items = action_items_by_id.get(action_id)
404 if items is None:
405 action_function = functions_by_id.get(action_id)
406 if action_function is None:
407 log.info('Action not found: %s', action_id)
408 action_items_by_id[action_id] = []
409 continue
411 function_module = functions_cache.get(action_id)
412 if function_module is None:
413 log.info('Failed to load action module: %s', action_id)
414 action_items_by_id[action_id] = []
415 continue
416 items = get_action_items_from_module(action_function, function_module)
417 action_items_by_id[action_id] = items
418 # Shallow copies keep per-model item dicts independent, as before
419 model['actions'].extend({**item} for item in items)
421 model['filters'] = []
422 for filter_id in filter_ids:
423 items = filter_items_by_id.get(filter_id)
424 if items is None:
425 filter_function = functions_by_id.get(filter_id)
426 if filter_function is None:
427 log.info('Filter not found: %s', filter_id)
428 filter_items_by_id[filter_id] = []
429 continue
431 function_module = functions_cache.get(filter_id)
432 if function_module is None:
433 log.info('Failed to load filter module: %s', filter_id)
434 filter_items_by_id[filter_id] = []
435 continue
436 if getattr(function_module, 'toggle', None):
437 items = get_filter_items_from_module(filter_function, function_module)
438 else:
439 items = []
440 filter_items_by_id[filter_id] = items
441 model['filters'].extend({**item} for item in items)
443 log.debug('get_all_models() returned %s models', len(models))
445 models_dict = {}
446 for model in models:
447 model = model.copy()
448 if model.get('ollama'):
449 # Keep the moving expiry in the API response, outside the registry signature.
450 model['ollama'] = model['ollama'].copy()
451 model['ollama'].pop('expires_at', None)
452 models_dict[model['id']] = model
453 if isinstance(request.app.state.MODELS, RedisDict):
454 try:
455 request.app.state.MODELS.set(models_dict)
456 except Exception as e:
457 log.warning(f'Failed to update Redis model cache, using in-process cache: {e}')
458 request.app.state.MODELS = models_dict
459 else:
460 request.app.state.MODELS = models_dict
462 return models
465async def check_model_access(user, model, model_info=None, db=None):
466 if model.get('arena'):
467 meta = model.get('info', {}).get('meta', {})
468 access_grants = meta.get('access_grants', [])
469 if not await has_access(
470 user.id,
471 permission='read',
472 access_grants=access_grants,
473 db=db,
474 ):
475 log.warning(
476 'Model access denied: user_id=%r model_id=%r reason=arena_read_denied',
477 user.id,
478 model.get('id'),
479 )
480 raise Exception('Model not found')
481 else:
482 # Callers that already fetched the row (chat completion entry) pass it in
483 if model_info is None or model_info.id != model.get('id'):
484 model_info = await Models.get_model_by_id(model.get('id'), db=db)
485 if not model_info:
486 log.warning(
487 'Model access denied: user_id=%r model_id=%r reason=model_unregistered',
488 user.id,
489 model.get('id'),
490 )
491 raise Exception('Model not found')
493 # One group-membership fetch shared by the direct check and every
494 # base-model hop; skipped when no check below needs it.
495 user_group_ids = None
496 if user.id != model_info.user_id or model_info.base_model_id:
497 user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
499 if not (
500 user.id == model_info.user_id
501 or await AccessGrants.has_access(
502 user_id=user.id,
503 resource_type='model',
504 resource_id=model_info.id,
505 permission='read',
506 user_group_ids=user_group_ids,
507 db=db,
508 )
509 ):
510 log.warning(
511 'Model access denied: user_id=%r model_id=%r reason=model_read_denied',
512 user.id,
513 model_info.id,
514 )
515 raise Exception('Model not found')
517 # Enforce access on chained base models
518 if not await has_base_model_access(
519 user.id, model_info, user_role=user.role, user_group_ids=user_group_ids, db=db
520 ):
521 raise Exception('Model not found')
524async def get_filtered_models(models, user, db=None):
525 # Filter out models that the user does not have access to
526 if ( 526 ↛ 529line 526 didn't jump to line 529 because the condition on line 526 was never true
527 user.role == 'user' or (user.role == 'admin' and not BYPASS_ADMIN_ACCESS_CONTROL)
528 ) and not BYPASS_MODEL_ACCESS_CONTROL:
529 model_infos = {}
530 for model in models:
531 if model.get('arena'):
532 continue
533 info = model.get('info')
534 if info:
535 model_infos[model['id']] = info
537 user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
539 # Batch-fetch accessible resource IDs in a single query instead of N has_access calls
540 accessible_model_ids = await AccessGrants.get_accessible_resource_ids(
541 user_id=user.id,
542 resource_type='model',
543 resource_ids=list(model_infos.keys()),
544 permission='read',
545 user_group_ids=user_group_ids,
546 db=db,
547 )
549 filtered_models = []
550 for model in models:
551 if model.get('arena'):
552 meta = model.get('info', {}).get('meta', {})
553 access_grants = meta.get('access_grants', [])
554 if await has_access(
555 user.id,
556 permission='read',
557 access_grants=access_grants,
558 user_group_ids=user_group_ids,
559 ):
560 filtered_models.append(model)
561 continue
563 model_info = model_infos.get(model['id'])
564 if model_info:
565 if (
566 (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
567 or user.id == model_info.get('user_id')
568 or model['id'] in accessible_model_ids
569 ):
570 filtered_models.append(model)
571 elif user.role == 'admin':
572 # No DB entry means no access control configured yet;
573 # only admins can see unconfigured models.
574 filtered_models.append(model)
576 return filtered_models
577 else:
578 return models