Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions backend/open_webui/routers/images.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
12 changes: 12 additions & 0 deletions backend/open_webui/routers/retrieval.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
19 changes: 19 additions & 0 deletions backend/open_webui/utils/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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]:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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", [])

Expand Down
17 changes: 17 additions & 0 deletions backend/open_webui/utils/permissions.py
Original file line number Diff line number Diff line change
@@ -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)
63 changes: 63 additions & 0 deletions tests/test_feature_permission_helper.py
Original file line number Diff line number Diff line change
@@ -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,
)