mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-22 16:44:06 -05:00
refac
This commit is contained in:
@@ -19,8 +19,8 @@ from fastapi import (
|
||||
from fastapi.responses import RedirectResponse, StreamingResponse
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING, PROFILE_IMAGE_ALLOWED_MIME_TYPES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
@@ -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
|
||||
|
||||
@@ -261,6 +262,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 +275,58 @@ 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 not request.app.state.MODELS:
|
||||
await get_all_models(request, user=user)
|
||||
for base_model in request.app.state.MODELS.values():
|
||||
base_model_id = base_model.get('id')
|
||||
if base_model.get('preset') or not base_model_id:
|
||||
continue
|
||||
|
||||
if form_data.id == base_model_id or (
|
||||
base_model.get('owned_by') == 'ollama' and form_data.id == base_model_id.split(':', 1)[0]
|
||||
):
|
||||
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 +412,12 @@ async def import_models(
|
||||
else:
|
||||
writable_model_ids = set(existing_model_ids)
|
||||
|
||||
base_model_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 +448,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(),
|
||||
@@ -451,6 +485,37 @@ async def import_models(
|
||||
model_data['meta'] = model_data.get('meta', {})
|
||||
model_data['params'] = model_data.get('params', {})
|
||||
new_model = ModelForm(**model_data)
|
||||
|
||||
if user.role != 'admin':
|
||||
if not new_model.base_model_id:
|
||||
log.warning(
|
||||
'import_models: user %s skipped model %s (no base model set)',
|
||||
user.id,
|
||||
model_id,
|
||||
)
|
||||
continue
|
||||
|
||||
if base_model_ids is None:
|
||||
base_model_ids = set()
|
||||
if not request.app.state.MODELS:
|
||||
await get_all_models(request, user=user)
|
||||
for base_model in request.app.state.MODELS.values():
|
||||
base_model_id = base_model.get('id')
|
||||
if base_model.get('preset') or not base_model_id:
|
||||
continue
|
||||
|
||||
base_model_ids.add(base_model_id)
|
||||
if base_model.get('owned_by') == 'ollama':
|
||||
base_model_ids.add(base_model_id.split(':', 1)[0])
|
||||
|
||||
if model_id in base_model_ids:
|
||||
log.warning(
|
||||
'import_models: user %s skipped model %s (id belongs to a base model)',
|
||||
user.id,
|
||||
model_id,
|
||||
)
|
||||
continue
|
||||
|
||||
new_model.access_grants = await filter_allowed_access_grants(
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
@@ -459,6 +524,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 +810,12 @@ 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
|
||||
|
||||
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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user