diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index e3d9e0e775..0fa7c7babf 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -147,6 +147,7 @@ from open_webui.models.config import Config from open_webui.models.functions import Functions from open_webui.models.messages import Messages from open_webui.models.models import Models, normalize_model_tags +from open_webui.models.groups import Groups, resolve_group_default_models from open_webui.models.users import Users from open_webui.routers import ( analytics, @@ -2277,6 +2278,24 @@ async def get_app_config(request: Request): if data is not None and 'id' in data and await is_valid_token(data, request.app.state.redis): user = await Users.get_user_by_id(data['id']) + group_defaults = None + if user is not None and user.role in ('admin', 'user'): + group_defaults, _ = resolve_group_default_models( + await Groups.get_groups_by_member_id(user.id, include_inherited=True) + ) + if group_defaults: + try: + models = (await get_models(request, user=user))['data'] + available = { + model['id'] + for model in models + if not ((model.get('info') or {}).get('meta') or {}).get('hidden', False) + } + group_defaults = [model_id for model_id in group_defaults if model_id in available] + except Exception: + log.exception('Unable to resolve available group default models') + group_defaults = None + onboarding = False if user is None: onboarding = not await Users.has_users() @@ -2425,7 +2444,7 @@ async def get_app_config(request: Request): }, **( { - 'default_models': config.get('ui.default_models'), + 'default_models': ','.join(group_defaults) if group_defaults else config.get('ui.default_models'), 'default_pinned_models': config.get('ui.default_pinned_models'), 'default_prompt_suggestions': config.get('ui.prompt_suggestions'), 'default_prompt_suggestions_i18n': config.get('ui.prompt_suggestions_i18n'), diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index aa1d51e5cc..6597928ce3 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -8,7 +8,7 @@ from open_webui.env import DEFAULT_GROUP_SHARE_PERMISSION from open_webui.internal.db import Base, JSONField, get_async_db_context, get_async_db from open_webui.models.access_grants import AccessGrant from open_webui.models.files import FileMetadataResponse -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, field_validator from sqlalchemy import ( JSON, BigInteger, @@ -129,6 +129,21 @@ class GroupForm(BaseModel): permissions: Optional[dict] = None data: Optional[dict] = None + @field_validator('data') + @classmethod + def validate_default_models(cls, data): + if data is None or 'config' not in data: + return data + config = data['config'] + if not isinstance(config, dict): + raise ValueError('Group config must be an object.') + if 'default_models' not in config or config['default_models'] is None: + return data + models = config['default_models'] + if not isinstance(models, list) or any(not isinstance(model, str) or not model.strip() for model in models): + raise ValueError('Default models must be a list of non-empty model IDs.') + return {**data, 'config': {**config, 'default_models': list(dict.fromkeys(model.strip() for model in models))}} + class UserIdsForm(BaseModel): user_ids: Optional[list[str]] = None @@ -149,6 +164,33 @@ class GroupHierarchyError(ValueError): self.status_code = status_code +def group_default_models(group): + return ((group.data or {}).get('config') or {}).get('default_models') or None + + +def resolve_group_default_models(groups): + """Resolve an ancestor-complete group list by depth, then creation time and ID.""" + by_id = {group.id: group for group in groups} + depths = {} + for group in groups: + path = [] + seen = set() + current = group + while current and current.id not in depths and current.id not in seen: + seen.add(current.id) + path.append(current.id) + current = by_id.get(current.parent_group_id) + depth = depths.get(current.id, -1) if current else -1 + for group_id in reversed(path): + depth += 1 + depths[group_id] = depth + configured = [group for group in groups if group_default_models(group)] + if not configured: + return None, None + winner = min(configured, key=lambda group: (-depths[group.id], group.created_at, group.id)) + return group_default_models(winner), winner.id + + def ancestor_groups(group_ids): """Identifier-only recursion: UNION also terminates on externally introduced cycles.""" chain = select(Group.id.label('group_id')).where(Group.id.in_(group_ids)).cte(recursive=True) @@ -504,13 +546,28 @@ class GroupTable: if 'parent_group_id' in form_data.model_fields_set: await validate_parent(session, id, form_data.parent_group_id) values['parent_group_id'] = form_data.parent_group_id + if 'data' in values: + old_data = group.data or {} + new_data = values['data'] + values['data'] = { + **old_data, + **new_data, + 'config': {**(old_data.get('config') or {}), **(new_data.get('config') or {})}, + } + defaults_changed = 'data' in values and ( + (values['data'].get('config') or {}).get('default_models') or None + ) != group_default_models(group) parent_changed = values.get('parent_group_id', group.parent_group_id) != group.parent_group_id if changes is not None: changes.update( old_parent_group_id=group.parent_group_id, parent_group_id=values.get('parent_group_id', group.parent_group_id), ) - if parent_changed or ('permissions' in values and values['permissions'] != group.permissions): + if ( + parent_changed + or defaults_changed + or ('permissions' in values and values['permissions'] != group.permissions) + ): members = group_user_memberships([id], True) affected = list((await session.execute(select(members.c.user_id))).scalars()) for key, value in values.items(): diff --git a/backend/open_webui/routers/groups.py b/backend/open_webui/routers/groups.py index 1701ea1c3e..8b4be1f3cb 100755 --- a/backend/open_webui/routers/groups.py +++ b/backend/open_webui/routers/groups.py @@ -12,6 +12,8 @@ from open_webui.internal.db import get_async_session from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import ( GroupForm, + group_default_models, + resolve_group_default_models, GroupHierarchyError, Group, GroupMember, @@ -423,8 +425,19 @@ async def preview_group_access( JSONCodec.loads(JSONCodec.dumps(inherited_permissions)), group.permissions or {} ) + default_models, source_group_id = resolve_group_default_models([group, *ancestors]) + if default_models is None: + default_models = [ + model.strip() for model in (await Config.get('ui.default_models') or '').split(',') if model.strip() + ] + return { 'group': {'id': group.id, 'name': group.name}, + 'default_models': { + 'local': group_default_models(group), + 'effective': default_models, + 'source_group_id': source_group_id, + }, 'models': { 'items': [{'id': m.id, 'name': m.name} for m in active_models if m.id in accessible_model_ids], 'total': len(active_models), diff --git a/src/lib/components/admin/Users/Groups.svelte b/src/lib/components/admin/Users/Groups.svelte index 6a79a3ecf1..53fbb1daef 100644 --- a/src/lib/components/admin/Users/Groups.svelte +++ b/src/lib/components/admin/Users/Groups.svelte @@ -1,5 +1,5 @@ @@ -132,6 +175,53 @@ +
+ {inheritedModelSource}{inheritedModelIds.length ? ': ' : ''}{inheritedModelIds + .map((id) => $models.find((model) => model.id === id)?.name || id) + .join(', ')} +
+ {/if} +