From 3e56261c5e6bba30bf8828551a0d6e976f811de4 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Wed, 11 Feb 2026 02:06:43 -0600 Subject: [PATCH] refac --- backend/open_webui/main.py | 15 +++ backend/open_webui/routers/tools.py | 32 +++-- backend/open_webui/utils/access_control.py | 110 +++++++++------- backend/open_webui/utils/db/access_control.py | 124 ------------------ backend/open_webui/utils/models.py | 16 +-- backend/open_webui/utils/tools.py | 5 +- src/lib/components/AddToolServerModal.svelte | 14 +- .../Evaluations/ArenaModelModal.svelte | 8 +- 8 files changed, 115 insertions(+), 209 deletions(-) delete mode 100644 backend/open_webui/utils/db/access_control.py diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 7d6ae9cba..72cb31fcf 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -815,6 +815,21 @@ app.state.config.ENABLE_USER_STATUS = ENABLE_USER_STATUS app.state.config.ENABLE_EVALUATION_ARENA_MODELS = ENABLE_EVALUATION_ARENA_MODELS app.state.config.EVALUATION_ARENA_MODELS = EVALUATION_ARENA_MODELS +# Migrate legacy access_control → access_grants on boot +from open_webui.utils.access_control import migrate_access_control + +connections = app.state.config.TOOL_SERVER_CONNECTIONS +if any("access_control" in c.get("config", {}) for c in connections): + for connection in connections: + migrate_access_control(connection.get("config", {})) + app.state.config.TOOL_SERVER_CONNECTIONS = connections + +arena_models = app.state.config.EVALUATION_ARENA_MODELS +if any("access_control" in m.get("meta", {}) for m in arena_models): + for model in arena_models: + migrate_access_control(model.get("meta", {})) + app.state.config.EVALUATION_ARENA_MODELS = arena_models + app.state.config.OAUTH_USERNAME_CLAIM = OAUTH_USERNAME_CLAIM app.state.config.OAUTH_PICTURE_CLAIM = OAUTH_PICTURE_CLAIM app.state.config.OAUTH_EMAIL_CLAIM = OAUTH_EMAIL_CLAIM diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 0d86272a2..ad9cab474 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -77,12 +77,21 @@ async def get_tools( ) # OpenAPI Tool Servers + server_access_grants = {} for server in await get_tool_servers(request): + connection = request.app.state.config.TOOL_SERVER_CONNECTIONS[ + server.get("idx", 0) + ] + server_config = connection.get("config", {}) + + server_id = f"server:{server.get('id')}" + server_access_grants[server_id] = server_config.get("access_grants", []) + tools.append( ToolUserResponse( **{ - "id": f"server:{server.get('id')}", - "user_id": f"server:{server.get('id')}", + "id": server_id, + "user_id": server_id, "name": server.get("openapi", {}) .get("info", {}) .get("title", "Tool Server"), @@ -91,11 +100,6 @@ async def get_tools( .get("info", {}) .get("description", ""), }, - "access_control": request.app.state.config.TOOL_SERVER_CONNECTIONS[ - server.get("idx", 0) - ] - .get("config", {}) - .get("access_control", None), "updated_at": int(time.time()), "created_at": int(time.time()), } @@ -119,20 +123,22 @@ async def get_tools( ) ) + server_config = server.get("config", {}) + + tool_id = f"server:mcp:{server.get('info', {}).get('id')}" + server_access_grants[tool_id] = server_config.get("access_grants", []) + tools.append( ToolUserResponse( **{ - "id": f"server:mcp:{server.get('info', {}).get('id')}", - "user_id": f"server:mcp:{server.get('info', {}).get('id')}", + "id": tool_id, + "user_id": tool_id, "name": server.get("info", {}).get("name", "MCP Tool Server"), "meta": { "description": server.get("info", {}).get( "description", "" ), }, - "access_control": server.get("config", {}).get( - "access_control", None - ), "updated_at": int(time.time()), "created_at": int(time.time()), **( @@ -161,7 +167,7 @@ async def get_tools( has_access( user.id, "read", - getattr(tool, "access_control", None), + server_access_grants.get(str(tool.id), []), user_group_ids, db=db, ) diff --git a/backend/open_webui/utils/access_control.py b/backend/open_webui/utils/access_control.py index 7784f6efd..3e9d02304 100644 --- a/backend/open_webui/utils/access_control.py +++ b/backend/open_webui/utils/access_control.py @@ -107,71 +107,79 @@ def has_permission( return get_permission(default_permissions, permission_hierarchy) -def get_permitted_group_and_user_ids( - type: str = "write", access_control: Optional[dict] = None -) -> Union[Dict[str, List[str]], None]: - if access_control is None: - return None - - permission_access = access_control.get(type, {}) - permitted_group_ids = permission_access.get("group_ids", []) - permitted_user_ids = permission_access.get("user_ids", []) - - return { - "group_ids": permitted_group_ids, - "user_ids": permitted_user_ids, - } - - def has_access( user_id: str, - type: str = "write", - access_control: Optional[dict] = None, + permission: str = "read", + access_grants: Optional[list] = None, user_group_ids: Optional[Set[str]] = None, - strict: bool = True, db: Optional[Any] = None, ) -> bool: - if access_control is None: - if strict: - return type == "read" - else: - return True + """ + Check if a user has the specified permission using an in-memory access_grants list. + + Used for config-driven resources (arena models, tool servers) that store + access control as JSON in PersistentConfig rather than in the access_grant DB table. + + Semantics: + - None or [] → private (owner-only, deny all) + - [{"principal_type": "user", "principal_id": "*", "permission": "read"}] → public read + - Specific grants → check user/group membership + """ + if not access_grants: + return False if user_group_ids is None: user_groups = Groups.get_groups_by_member_id(user_id, db=db) user_group_ids = {group.id for group in user_groups} - permitted_ids = get_permitted_group_and_user_ids(type, access_control) - if permitted_ids is None: - return False - - permitted_group_ids = permitted_ids.get("group_ids", []) - permitted_user_ids = permitted_ids.get("user_ids", []) - - return user_id in permitted_user_ids or any( - group_id in permitted_group_ids for group_id in user_group_ids - ) + for grant in access_grants: + if not isinstance(grant, dict): + continue + if grant.get("permission") != permission: + continue + principal_type = grant.get("principal_type") + principal_id = grant.get("principal_id") + if principal_type == "user" and (principal_id == "*" or principal_id == user_id): + return True + if principal_type == "group" and user_group_ids and principal_id in user_group_ids: + return True -# Get all users with access to a resource -def get_users_with_access( - type: str = "write", access_control: Optional[dict] = None, db: Optional[Any] = None -) -> list[UserModel]: - if access_control is None: - result = Users.get_users(filter={"roles": ["!pending"]}, db=db) - return result.get("users", []) + return False - permitted_ids = get_permitted_group_and_user_ids(type, access_control) - if permitted_ids is None: - return [] - permitted_group_ids = permitted_ids.get("group_ids", []) - permitted_user_ids = permitted_ids.get("user_ids", []) +def migrate_access_control(data: dict, ac_key: str = "access_control", grants_key: str = "access_grants") -> None: + """ + Auto-migrate a config dict in-place from legacy access_control dict to access_grants list. - user_ids_with_access = set(permitted_user_ids) + If `grants_key` already exists, does nothing. + If `ac_key` exists (old format), converts it and stores as `grants_key`, then removes `ac_key`. + """ + if grants_key in data: + return - group_user_ids_map = Groups.get_group_user_ids_by_ids(permitted_group_ids, db=db) - for user_ids in group_user_ids_map.values(): - user_ids_with_access.update(user_ids) + access_control = data.get(ac_key) + if access_control is None and ac_key not in data: + return - return Users.get_users_by_user_ids(list(user_ids_with_access), db=db) + grants: List[Dict[str, str]] = [] + if access_control and isinstance(access_control, dict): + for perm in ["read", "write"]: + perm_data = access_control.get(perm, {}) + if not perm_data: + continue + for group_id in perm_data.get("group_ids", []): + grants.append({ + "principal_type": "group", + "principal_id": group_id, + "permission": perm, + }) + for uid in perm_data.get("user_ids", []): + grants.append({ + "principal_type": "user", + "principal_id": uid, + "permission": perm, + }) + + data[grants_key] = grants + data.pop(ac_key, None) diff --git a/backend/open_webui/utils/db/access_control.py b/backend/open_webui/utils/db/access_control.py deleted file mode 100644 index 75bd337f8..000000000 --- a/backend/open_webui/utils/db/access_control.py +++ /dev/null @@ -1,124 +0,0 @@ -from pydantic import BaseModel, ConfigDict -from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON -from sqlalchemy.dialects.postgresql import JSONB - - -from sqlalchemy import or_, func, select, and_, text, cast, or_, and_, func - - -def has_permission(db, DocumentModel, query, filter: dict, permission: str = "read"): - group_ids = filter.get("group_ids", []) - user_id = filter.get("user_id") - dialect_name = db.bind.dialect.name - - conditions = [] - - # Handle read_only permission separately - if permission == "read_only": - # For read_only, we want items where: - # 1. User has explicit read permission (via groups or user-level) - # 2. BUT does NOT have write permission - # 3. Public items are NOT considered read_only - - read_conditions = [] - - # Group-level read permission - if group_ids: - group_read_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_read_conditions.append( - DocumentModel.access_control["read"]["group_ids"].contains(gid) - ) - elif dialect_name == "postgresql": - group_read_conditions.append( - cast( - DocumentModel.access_control["read"]["group_ids"], - JSONB, - ).contains([gid]) - ) - - if group_read_conditions: - read_conditions.append(or_(*group_read_conditions)) - - # Combine read conditions - if read_conditions: - has_read = or_(*read_conditions) - else: - # If no read conditions, return empty result - return query.filter(False) - - # Now exclude items where user has write permission - write_exclusions = [] - - # Exclude items owned by user (they have implicit write) - if user_id: - write_exclusions.append(DocumentModel.user_id != user_id) - - # Exclude items where user has explicit write permission via groups - if group_ids: - group_write_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_write_conditions.append( - DocumentModel.access_control["write"]["group_ids"].contains(gid) - ) - elif dialect_name == "postgresql": - group_write_conditions.append( - cast( - DocumentModel.access_control["write"]["group_ids"], - JSONB, - ).contains([gid]) - ) - - if group_write_conditions: - # User should NOT have write permission - write_exclusions.append(~or_(*group_write_conditions)) - - # Exclude public items (items without access_control) - write_exclusions.append(DocumentModel.access_control.isnot(None)) - write_exclusions.append(cast(DocumentModel.access_control, String) != "null") - - # Combine: has read AND does not have write AND not public - if write_exclusions: - query = query.filter(and_(has_read, *write_exclusions)) - else: - query = query.filter(has_read) - - return query - - # Original logic for other permissions (read, write, etc.) - # Public access conditions - if group_ids or user_id: - conditions.extend( - [ - DocumentModel.access_control.is_(None), - cast(DocumentModel.access_control, String) == "null", - ] - ) - - # User-level permission (owner has all permissions) - if user_id: - conditions.append(DocumentModel.user_id == user_id) - - # Group-level permission - if group_ids: - group_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_conditions.append( - DocumentModel.access_control[permission]["group_ids"].contains(gid) - ) - elif dialect_name == "postgresql": - group_conditions.append( - cast( - DocumentModel.access_control[permission]["group_ids"], - JSONB, - ).contains([gid]) - ) - conditions.append(or_(*group_conditions)) - - if conditions: - query = query.filter(or_(*conditions)) - - return query diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index 4224605f1..986c0c601 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -340,12 +340,12 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) def check_model_access(user, model, db=None): if model.get("arena"): + meta = model.get("info", {}).get("meta", {}) + access_grants = meta.get("access_grants", []) if not has_access( user.id, - type="read", - access_control=model.get("info", {}) - .get("meta", {}) - .get("access_control", {}), + permission="read", + access_grants=access_grants, db=db, ): raise Exception("Model not found") @@ -384,12 +384,12 @@ def get_filtered_models(models, user, db=None): } for model in models: if model.get("arena"): + meta = model.get("info", {}).get("meta", {}) + access_grants = meta.get("access_grants", []) if has_access( user.id, - type="read", - access_control=model.get("info", {}) - .get("meta", {}) - .get("access_control", {}), + permission="read", + access_grants=access_grants, user_group_ids=user_group_ids, ): filtered_models.append(model) diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 5bb523f83..88479f40d 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -149,8 +149,9 @@ def has_tool_server_access( if user_group_ids is None: user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - access_control = server_connection.get("config", {}).get("access_control", None) - return has_access(user.id, "read", access_control, user_group_ids) + server_config = server_connection.get("config", {}) + access_grants = server_config.get("access_grants", []) + return has_access(user.id, "read", access_grants, user_group_ids) async def get_tools( diff --git a/src/lib/components/AddToolServerModal.svelte b/src/lib/components/AddToolServerModal.svelte index 764e2259b..5b64b692a 100644 --- a/src/lib/components/AddToolServerModal.svelte +++ b/src/lib/components/AddToolServerModal.svelte @@ -48,7 +48,7 @@ let headers = ''; let functionNameFilterList = ''; - let accessControl = {}; + let accessGrants = []; let id = ''; let name = ''; @@ -149,7 +149,7 @@ key, config: { enable: enable, - access_control: accessControl + access_grants: accessGrants }, info: { id, @@ -206,7 +206,7 @@ if (data.config) { enable = data.config.enable ?? true; - accessControl = data.config.access_control ?? {}; + accessGrants = data.config.access_grants ?? []; } toast.success($i18n.t('Import successful')); @@ -305,7 +305,7 @@ config: { enable: enable, function_name_filter_list: functionNameFilterList, - access_control: accessControl + access_grants: accessGrants }, info: { id: id, @@ -339,7 +339,7 @@ enable = true; functionNameFilterList = ''; - accessControl = null; + accessGrants = []; }; const init = () => { @@ -363,7 +363,7 @@ enable = connection.config?.enable ?? true; functionNameFilterList = connection.config?.function_name_filter_list ?? ''; - accessControl = connection.config?.access_control ?? null; + accessGrants = connection.config?.access_grants ?? []; } }; @@ -819,7 +819,7 @@
- +
{/if} diff --git a/src/lib/components/admin/Settings/Evaluations/ArenaModelModal.svelte b/src/lib/components/admin/Settings/Evaluations/ArenaModelModal.svelte index ef78d2352..2726dd13b 100644 --- a/src/lib/components/admin/Settings/Evaluations/ArenaModelModal.svelte +++ b/src/lib/components/admin/Settings/Evaluations/ArenaModelModal.svelte @@ -44,7 +44,7 @@ let modelIds = []; let filterMode = 'include'; - let accessControl = {}; + let accessGrants = []; let imageInputElement; let loading = false; @@ -83,7 +83,7 @@ description: description || null, model_ids: modelIds.length > 0 ? modelIds : null, filter_mode: modelIds.length > 0 ? (filterMode ? filterMode : null) : null, - access_control: accessControl + access_grants: accessGrants } }; @@ -107,7 +107,7 @@ description = model.meta.description; modelIds = model.meta.model_ids || []; filterMode = model.meta?.filter_mode ?? 'include'; - accessControl = 'access_control' in model.meta ? model.meta.access_control : {}; + accessGrants = model.meta.access_grants ?? []; } }; @@ -293,7 +293,7 @@
- +