diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index 0fab100c46..9acd0c9b70 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -229,10 +229,25 @@ class ModelsTable: ) return models - async def get_base_models(self, db: AsyncSession | None = None) -> list[ModelModel]: + @staticmethod + def _meta_has_tag(meta: dict | None, tag: str) -> bool: + if not meta: + return False + + for raw_tag in meta.get('tags', []): + name = raw_tag.get('name') if isinstance(raw_tag, dict) else str(raw_tag) + if name == tag: + return True + + return False + + async def get_base_models(self, tag: str | None = None, db: AsyncSession | None = None) -> list[ModelModel]: async with get_async_db_context(db) as db: - result = await db.execute(select(Model).filter(Model.base_model_id == None)) + result = await db.execute(select(Model).filter(Model.base_model_id.is_(None))) all_models = result.scalars().all() + if tag: + all_models = [model for model in all_models if self._meta_has_tag(model.meta, tag)] + model_ids = [model.id for model in all_models] grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) return [ @@ -395,11 +410,14 @@ class ModelsTable: self, user_id: str, is_admin: bool = False, + is_base_model: bool = False, db: AsyncSession | None = None, ) -> set[str]: """Extract unique tag names from model meta, querying only the meta column.""" async with get_async_db_context(db) as db: - stmt = select(Model.meta).filter(Model.base_model_id != None) + stmt = select(Model.meta).filter( + Model.base_model_id.is_(None) if is_base_model else Model.base_model_id.is_not(None) + ) if not is_admin: user_groups = await Groups.get_groups_by_member_id(user_id, db=db) diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index af16958e8e..2af08b5b20 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -198,9 +198,19 @@ async def get_models( ########################### +@router.get('/base/tags', response_model=list[str]) +async def get_base_model_tags(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + tags = await Models.get_all_tags(user_id=user.id, is_admin=True, is_base_model=True, db=db) + return sorted(tags) + + @router.get('/base', response_model=list[ModelResponse]) -async def get_base_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): - return await Models.get_base_models(db=db) +async def get_base_models( + tag: str | None = None, + user=Depends(get_admin_user), + db: AsyncSession = Depends(get_async_session), +): + return await Models.get_base_models(tag=tag, db=db) ########################### diff --git a/src/lib/apis/models/index.ts b/src/lib/apis/models/index.ts index e7abaa309e..11866e1fb6 100644 --- a/src/lib/apis/models/index.ts +++ b/src/lib/apis/models/index.ts @@ -90,6 +90,37 @@ export const getModelTags = async (token: string = '') => { return res; }; +export const getBaseModelTags = async (token: string = '') => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/models/base/tags`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .then((json) => { + return json; + }) + .catch((err) => { + error = err; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const importModels = async (token: string, models: object[]) => { let error = null; @@ -118,10 +149,15 @@ export const importModels = async (token: string, models: object[]) => { return res; }; -export const getBaseModels = async (token: string = '') => { +export const getBaseModels = async (token: string = '', tag: string = '') => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/models/base`, { + const searchParams = new URLSearchParams(); + if (tag) { + searchParams.append('tag', tag); + } + + const res = await fetch(`${WEBUI_API_BASE_URL}/models/base?${searchParams.toString()}`, { method: 'GET', headers: { Accept: 'application/json', diff --git a/src/lib/components/admin/Settings/Models.svelte b/src/lib/components/admin/Settings/Models.svelte index cacb5b5e73..6ad5acb260 100644 --- a/src/lib/components/admin/Settings/Models.svelte +++ b/src/lib/components/admin/Settings/Models.svelte @@ -10,6 +10,7 @@ import { createNewModel, deleteAllModels, + getBaseModelTags, getBaseModels, toggleModelById, updateModelById, @@ -46,8 +47,11 @@ import Dropdown from '$lib/components/common/Dropdown.svelte'; import AdminViewSelector from './Models/AdminViewSelector.svelte'; + import TagSelector from '$lib/components/workspace/common/TagSelector.svelte'; import Pagination from '$lib/components/common/Pagination.svelte'; + type ModelListItem = { id: string; name?: string }; + let shiftKey = false; let modelsImportInProgress = false; @@ -56,8 +60,8 @@ let models = null; - let workspaceModels = null; - let baseModels = null; + let workspaceModels: ModelListItem[] = []; + let baseModels: ModelListItem[] = []; let filteredModels = []; let selectedModelId = null; @@ -66,6 +70,8 @@ let showManageModal = false; let viewOption = ''; // '' = All, 'enabled', 'disabled', 'visible', 'hidden' + let tags: string[] = []; + let selectedTag = ''; const perPage = 30; let currentPage = 1; @@ -175,27 +181,35 @@ const init = async () => { models = null; - workspaceModels = await getBaseModels(localStorage.token); + tags = await getBaseModelTags(localStorage.token); + if (selectedTag && !tags.includes(selectedTag)) { + selectedTag = ''; + } + + workspaceModels = await getBaseModels(localStorage.token, selectedTag); baseModels = await getModels(localStorage.token, null, true); + const workspaceModelIds = new Set(workspaceModels.map((wm: ModelListItem) => wm.id)); - models = baseModels.map((m) => { - const workspaceModel = workspaceModels.find((wm) => wm.id === m.id); + models = baseModels + .filter((m: ModelListItem) => !selectedTag || workspaceModelIds.has(m.id)) + .map((m: ModelListItem) => { + const workspaceModel = workspaceModels.find((wm: ModelListItem) => wm.id === m.id); - if (workspaceModel) { - return { - ...m, - ...workspaceModel - }; - } else { - return { - ...m, - id: m.id, - name: m.name, + if (workspaceModel) { + return { + ...m, + ...workspaceModel + }; + } else { + return { + ...m, + id: m.id, + name: m.name, - is_active: true - }; - } - }); + is_active: true + }; + } + }); _models.set( await getModels( @@ -501,6 +515,16 @@ class="flex gap-0.5 w-fit text-center text-sm rounded-full bg-transparent whitespace-nowrap" > + {#if (tags ?? []).length > 0} + ({ value: tag, label: tag }))} + onChange={async () => { + currentPage = 1; + await init(); + }} + /> + {/if}