refac
This commit is contained in:
@@ -34,8 +34,10 @@ from open_webui.utils.plugin import (
|
||||
load_function_module_by_id,
|
||||
get_function_module_from_cache,
|
||||
)
|
||||
from open_webui.utils.access_control import check_model_access
|
||||
|
||||
from open_webui.env import GLOBAL_LOG_LEVEL
|
||||
from open_webui.env import GLOBAL_LOG_LEVEL, BYPASS_MODEL_ACCESS_CONTROL
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
|
||||
from open_webui.utils.misc import (
|
||||
add_or_update_system_message,
|
||||
@@ -260,6 +262,10 @@ async def generate_function_chat_completion(request, form_data, user, models: di
|
||||
if model_info.base_model_id:
|
||||
form_data['model'] = model_info.base_model_id
|
||||
|
||||
if not BYPASS_MODEL_ACCESS_CONTROL:
|
||||
bypass = isinstance(user, UserModel) and user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL
|
||||
await check_model_access(user if isinstance(user, UserModel) else UserModel(**user), model_info, bypass)
|
||||
|
||||
params = model_info.params.model_dump()
|
||||
|
||||
if params:
|
||||
|
||||
@@ -1823,6 +1823,8 @@ async def chat_completion(
|
||||
request.state.metadata = metadata
|
||||
form_data['metadata'] = metadata
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.debug(f'Error processing chat metadata: {e}')
|
||||
raise HTTPException(
|
||||
|
||||
@@ -1383,29 +1383,9 @@ async def generate_responses(
|
||||
if model_info.base_model_id:
|
||||
payload['model'] = model_info.base_model_id
|
||||
|
||||
# Check if user has access to the model
|
||||
if user.role == 'user':
|
||||
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
|
||||
if not (
|
||||
user.id == model_info.user_id
|
||||
or await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='model',
|
||||
resource_id=model_info.id,
|
||||
permission='read',
|
||||
user_group_ids=user_group_ids,
|
||||
)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
|
||||
)
|
||||
await check_model_access(user, model_info)
|
||||
else:
|
||||
if user.role != 'admin':
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
|
||||
)
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx)
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
|
||||
@@ -339,10 +339,11 @@ async def check_model_access(
|
||||
raise HTTPException(status_code=403, detail='Model not found')
|
||||
|
||||
# Enforce access on chained base models
|
||||
if not await has_base_model_chain_access(
|
||||
if not await has_base_model_access(
|
||||
user.id, model_info, user_group_ids=user_group_ids
|
||||
):
|
||||
raise HTTPException(status_code=403, detail='Model not found')
|
||||
else:
|
||||
if user.role != 'admin':
|
||||
raise HTTPException(status_code=403, detail='Model not found')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user