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

1import asyncio 

2import copy 

3import logging 

4import sys 

5 

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) 

28 

29logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) 

30log = logging.getLogger(__name__) 

31 

32BASE_MODELS_CACHE_KEY = f'{REDIS_KEY_PREFIX}:models:base' 

33 

34 

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 ] 

51 

52 

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'] 

56 

57 

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) 

63 

64 openai_models, ollama_models, function_models = await asyncio.gather(openai_task, ollama_task, function_task) 

65 

66 return function_models + openai_models + ollama_models 

67 

68 

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 = [] 

83 

84 redis = getattr(request.app.state, 'redis', None) 

85 use_cache = config.get('models.base_models_cache') and not refresh 

86 base_models = None 

87 

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 

98 

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 

107 

108 # deep copy the base models to avoid modifying the original list 

109 models = [model.copy() for model in base_models] 

110 

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 [] 

114 

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 

150 

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} 

157 

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() 

166 

167 custom_models = await Models.get_all_models() 

168 

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 

175 

176 existing_ids = {m['id'] for m in models} 

177 

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) 

182 

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) 

192 

193 action_ids = [] 

194 filter_ids = [] 

195 

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', [])) 

201 

202 if 'params' in model['info']: 

203 del model['info']['params'] 

204 

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] 

209 

210 elif custom_model.is_active: 

211 if custom_model.id in existing_ids: 

212 continue 

213 

214 owned_by = 'openai' 

215 connection_type = None 

216 pipe = None 

217 

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) 

226 

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 } 

239 

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'] 

249 

250 model['info'] = info 

251 

252 action_ids = [] 

253 filter_ids = [] 

254 

255 if custom_model.meta: 

256 meta = custom_model.meta.model_dump() 

257 

258 if ENABLE_PLUGINS and 'actionIds' in meta: 

259 action_ids.extend(meta['actionIds']) 

260 

261 if ENABLE_PLUGINS and 'filterIds' in meta: 

262 filter_ids.extend(meta['filterIds']) 

263 

264 model['action_ids'] = action_ids 

265 model['filter_ids'] = filter_ids 

266 

267 models.append(model) 

268 

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 ] 

299 

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 ] 

313 

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) 

321 

322 functions_by_id = {f.id: f for f in await Functions.get_functions_by_ids(list(all_function_ids))} 

323 

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) 

336 

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 {} 

340 

341 if default_metadata: 

342 for model in models: 

343 info = model.get('info') 

344 

345 if info is None: 

346 model['info'] = {'meta': copy.deepcopy(default_metadata)} 

347 continue 

348 

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) 

357 

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) 

362 

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 = {} 

366 

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 

381 

382 action_items_by_id = {} 

383 filter_items_by_id = {} 

384 

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)) 

392 

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() 

400 

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 

410 

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) 

420 

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 

430 

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) 

442 

443 log.debug('get_all_models() returned %s models', len(models)) 

444 

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 

461 

462 return models 

463 

464 

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') 

492 

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)} 

498 

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') 

516 

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') 

522 

523 

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 

536 

537 user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} 

538 

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 ) 

548 

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 

562 

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) 

575 

576 return filtered_models 

577 else: 

578 return models