diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index 808b5b431e..191b270e45 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -40,6 +40,7 @@ from open_webui.utils.access_control import filter_allowed_access_grants, has_pe from open_webui.utils.access_control.files import has_access_to_file from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.chat_variables import get_chat_variables_schema +from open_webui.utils.models import get_all_models from pydantic import BaseModel from sqlalchemy.ext.asyncio import AsyncSession @@ -91,6 +92,21 @@ def is_valid_model_id(model_id: str) -> bool: return model_id and len(model_id) <= 256 +async def _reserved_base_model_ids(request: Request, user) -> set[str]: + """Ids that already answer to a base model, keyed the same way get_all_models keys its base model lookup.""" + if not request.app.state.MODELS: + await get_all_models(request, user=user) + + ids = set() + for model in request.app.state.MODELS.values(): + if model.get('preset'): + continue + ids.add(model['id']) + if model.get('owned_by') == 'ollama': + ids.add(model['id'].split(':')[0]) + return ids + + async def _verify_knowledge_file_access( knowledge_items: list | None, user, @@ -261,6 +277,12 @@ async def create_new_model( detail=ERROR_MESSAGES.UNAUTHORIZED, ) + if not is_valid_model_id(form_data.id): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.MODEL_ID_TOO_LONG, + ) + model = await Models.get_model_by_id(form_data.id, db=db) if model: raise HTTPException( @@ -268,43 +290,49 @@ async def create_new_model( detail=ERROR_MESSAGES.MODEL_ID_TAKEN, ) - if not is_valid_model_id(form_data.id): - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=ERROR_MESSAGES.MODEL_ID_TOO_LONG, - ) - - else: - await _verify_knowledge_file_access( - getattr(form_data.meta, 'knowledge', None) if form_data.meta else None, - user, - db, - ) - - form_data.access_grants = await filter_allowed_access_grants( - await Config.get('user.permissions'), - user.id, - user.role, - form_data.access_grants, - 'sharing.public_models', - ) - - model = await Models.insert_new_model(form_data, user.id, db=db) - if model: - await publish_event( - request, - EVENTS.MODEL_CREATED, - actor=user, - subject_id=model.id, - data={'name': model.name}, - ) - return model - else: + if user.role != 'admin': + if not form_data.base_model_id: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, - detail=ERROR_MESSAGES.DEFAULT(), + detail=ERROR_MESSAGES.UNAUTHORIZED, ) + if form_data.id in await _reserved_base_model_ids(request, user): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.MODEL_ID_TAKEN, + ) + + await _verify_knowledge_file_access( + getattr(form_data.meta, 'knowledge', None) if form_data.meta else None, + user, + db, + ) + + form_data.access_grants = await filter_allowed_access_grants( + await Config.get('user.permissions'), + user.id, + user.role, + form_data.access_grants, + 'sharing.public_models', + ) + + model = await Models.insert_new_model(form_data, user.id, db=db) + if not model: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.DEFAULT(), + ) + + await publish_event( + request, + EVENTS.MODEL_CREATED, + actor=user, + subject_id=model.id, + data={'name': model.name}, + ) + return model + ############################ # ExportModels @@ -390,12 +418,14 @@ async def import_models( else: writable_model_ids = set(existing_model_ids) + # Filled on the first insert that needs it, since filling it can fan out to every provider. + reserved_ids = None + imported_ids = [] for model_data in data: model_id = model_data.get('id') if model_id and is_valid_model_id(model_id): - imported_ids.append(model_id) # Defense-in-depth: skip models referencing inaccessible files try: await _verify_knowledge_file_access( @@ -426,6 +456,18 @@ async def import_models( ) continue + if ( + user.role != 'admin' + and existing_model.base_model_id + and not model_data.get('base_model_id', existing_model.base_model_id) + ): + log.warning( + 'import_models: user %s skipped model %s (cannot clear base model)', + user.id, + model_id, + ) + continue + # Update existing model model_data['meta'] = { **existing_model.meta.model_dump(), @@ -447,6 +489,26 @@ async def import_models( ) await Models.update_model_by_id(model_id, updated_model, db=db) else: + if user.role != 'admin': + if not model_data.get('base_model_id'): + log.warning( + 'import_models: user %s skipped model %s (no base model set)', + user.id, + model_id, + ) + continue + + if reserved_ids is None: + reserved_ids = await _reserved_base_model_ids(request, user) + + if model_id in reserved_ids: + log.warning( + 'import_models: user %s skipped model %s (id is reserved by a base model)', + user.id, + model_id, + ) + continue + # Insert new model model_data['meta'] = model_data.get('meta', {}) model_data['params'] = model_data.get('params', {}) @@ -459,6 +521,8 @@ async def import_models( 'sharing.public_models', ) await Models.insert_new_model(user_id=user.id, form_data=new_model, db=db) + + imported_ids.append(model_id) await publish_event( request, EVENTS.MODEL_IMPORTED, @@ -743,6 +807,13 @@ async def update_model_by_id( if 'base_model_id' not in form_data.model_fields_set: form_data.base_model_id = model.base_model_id + # Clearing the base model turns the row into an override of the model sharing its id. + if user.role != 'admin' and model.base_model_id and not form_data.base_model_id: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.UNAUTHORIZED, + ) + if 'profile_image_url' not in form_data.meta.model_fields_set: form_data.meta.profile_image_url = model.meta.profile_image_url