diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index 52686a5841b0..51c31ed50d9d 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -14,7 +14,9 @@ from open_webui.constants import ERROR_MESSAGES from open_webui.env import ENABLE_FORWARD_USER_INFO_HEADERS, SRC_LOG_LEVELS from open_webui.routers.files import upload_file +from open_webui.utils.access_control import has_permission from open_webui.utils.auth import get_admin_user, get_verified_user +from open_webui.utils.permissions import user_has_permission from open_webui.utils.images.comfyui import ( ComfyUIGenerateImageForm, ComfyUIWorkflow, @@ -471,6 +473,17 @@ async def image_generations( form_data: GenerateImageForm, user=Depends(get_verified_user), ): + if not user_has_permission( + user, + "features.image_generation", + request.app.state.config.USER_PERMISSIONS, + has_permission, + ): + raise HTTPException( + status_code=403, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + width, height = tuple(map(int, request.app.state.config.IMAGE_SIZE.split("x"))) r = None diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index ee6f99fbb5ba..4c08ff423844 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -78,7 +78,9 @@ from open_webui.utils.misc import ( calculate_sha256_string, ) +from open_webui.utils.access_control import has_permission from open_webui.utils.auth import get_admin_user, get_verified_user +from open_webui.utils.permissions import user_has_permission from open_webui.config import ( ENV, @@ -1833,6 +1835,16 @@ def search_web(request: Request, engine: str, query: str) -> list[SearchResult]: async def process_web_search( request: Request, form_data: SearchForm, user=Depends(get_verified_user) ): + if not user_has_permission( + user, + "features.web_search", + request.app.state.config.USER_PERMISSIONS, + has_permission, + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) urls = [] try: diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index b1e69db2640f..dfceb2848ddc 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -76,6 +76,8 @@ convert_logit_bias_input_to_json, ) from open_webui.utils.tools import get_tools +from open_webui.utils.access_control import has_permission +from open_webui.utils.permissions import user_has_permission from open_webui.utils.plugin import load_function_module_by_id from open_webui.utils.filter import ( get_sorted_filter_ids, @@ -104,6 +106,19 @@ log.setLevel(SRC_LOG_LEVELS["MAIN"]) +def require_feature_permission(request: Request, user: UserModel, feature: str) -> None: + if not user_has_permission( + user, + f"features.{feature}", + request.app.state.config.USER_PERMISSIONS, + has_permission, + ): + raise HTTPException( + status_code=403, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + async def chat_completion_tools_handler( request: Request, body: dict, extra_params: dict, user: UserModel, models, tools ) -> tuple[dict, dict]: @@ -830,16 +845,19 @@ async def process_chat_payload(request, form_data, user, metadata, model): ) if "web_search" in features and features["web_search"]: + require_feature_permission(request, user, "web_search") form_data = await chat_web_search_handler( request, form_data, extra_params, user ) if "image_generation" in features and features["image_generation"]: + require_feature_permission(request, user, "image_generation") form_data = await chat_image_generation_handler( request, form_data, extra_params, user ) if "code_interpreter" in features and features["code_interpreter"]: + require_feature_permission(request, user, "code_interpreter") form_data["messages"] = add_or_update_user_message( ( request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE @@ -887,6 +905,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): ) if tool_servers: + require_feature_permission(request, user, "direct_tool_servers") for tool_server in tool_servers: tool_specs = tool_server.pop("specs", []) diff --git a/backend/open_webui/utils/permissions.py b/backend/open_webui/utils/permissions.py new file mode 100644 index 000000000000..e61a1c88b014 --- /dev/null +++ b/backend/open_webui/utils/permissions.py @@ -0,0 +1,17 @@ +from typing import Any, Callable, Dict + + +def user_has_permission( + user: Any, + permission_key: str, + default_permissions: Dict[str, Any], + permission_checker: Callable[[str, str, Dict[str, Any]], bool], +) -> bool: + if getattr(user, "role", None) == "admin": + return True + + user_id = getattr(user, "id", None) + if not user_id: + return False + + return permission_checker(user_id, permission_key, default_permissions) diff --git a/tests/test_feature_permission_helper.py b/tests/test_feature_permission_helper.py new file mode 100644 index 000000000000..cf5982a9f6ef --- /dev/null +++ b/tests/test_feature_permission_helper.py @@ -0,0 +1,63 @@ +import importlib.util +from pathlib import Path +from types import SimpleNamespace + + +def _load_permissions_module(): + module_path = ( + Path(__file__).resolve().parents[1] + / "backend" + / "open_webui" + / "utils" + / "permissions.py" + ) + spec = importlib.util.spec_from_file_location("permissions", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_admin_users_do_not_call_permission_checker(): + permissions = _load_permissions_module() + + def deny(*args): + raise AssertionError("admin users should bypass delegated permission checks") + + assert permissions.user_has_permission( + SimpleNamespace(id="admin-user", role="admin"), + "features.code_interpreter", + {"features": {"code_interpreter": False}}, + deny, + ) + + +def test_regular_users_delegate_to_permission_checker(): + permissions = _load_permissions_module() + calls = [] + + def allow(user_id, permission_key, default_permissions): + calls.append((user_id, permission_key, default_permissions)) + return permission_key == "features.web_search" + + default_permissions = {"features": {"web_search": True}} + + assert permissions.user_has_permission( + SimpleNamespace(id="regular-user", role="user"), + "features.web_search", + default_permissions, + allow, + ) + assert calls == [ + ("regular-user", "features.web_search", default_permissions), + ] + + +def test_users_without_ids_are_denied(): + permissions = _load_permissions_module() + + assert not permissions.user_has_permission( + SimpleNamespace(role="user"), + "features.direct_tool_servers", + {"features": {"direct_tool_servers": True}}, + lambda *args: True, + )