diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index b2bccc195..eff83a54d 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -481,6 +481,52 @@ async def update_model_by_id( return model +############################ +# UpdateModelAccessById +############################ + + +class ModelAccessGrantsForm(BaseModel): + id: str + access_grants: list[dict] + + +@router.post("/model/access/update", response_model=Optional[ModelModel]) +async def update_model_access_by_id( + form_data: ModelAccessGrantsForm, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + model = Models.get_model_by_id(form_data.id, db=db) + if not model: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + if ( + model.user_id != user.id + and not AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="write", + db=db, + ) + and user.role != "admin" + ): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + AccessGrants.set_access_grants( + "model", form_data.id, form_data.access_grants, db=db + ) + + return Models.get_model_by_id(form_data.id, db=db) + + ############################ # DeleteModelById ############################ diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 015bde232..0d86272a2 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -510,6 +510,52 @@ async def update_tools_by_id( ) +############################ +# UpdateToolAccessById +############################ + + +class ToolAccessGrantsForm(BaseModel): + access_grants: list[dict] + + +@router.post("/id/{id}/access/update", response_model=Optional[ToolModel]) +async def update_tool_access_by_id( + id: str, + form_data: ToolAccessGrantsForm, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + tools = Tools.get_tool_by_id(id, db=db) + if not tools: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + if ( + tools.user_id != user.id + and not AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tools.id, + permission="write", + db=db, + ) + and user.role != "admin" + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.UNAUTHORIZED, + ) + + AccessGrants.set_access_grants( + "tool", id, form_data.access_grants, db=db + ) + + return Tools.get_tool_by_id(id, db=db) + + ############################ # DeleteToolsById ############################ diff --git a/src/lib/apis/models/index.ts b/src/lib/apis/models/index.ts index d03a83e9c..0151a55c1 100644 --- a/src/lib/apis/models/index.ts +++ b/src/lib/apis/models/index.ts @@ -281,6 +281,39 @@ export const updateModelById = async (token: string, id: string, model: object) return res; }; +export const updateModelAccessGrants = async ( + token: string, + id: string, + accessGrants: any[] +) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/models/model/access/update`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify({ id, access_grants: accessGrants }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const deleteModelById = async (token: string, id: string) => { let error = null; diff --git a/src/lib/apis/tools/index.ts b/src/lib/apis/tools/index.ts index 2038e46ac..103e05f39 100644 --- a/src/lib/apis/tools/index.ts +++ b/src/lib/apis/tools/index.ts @@ -225,6 +225,39 @@ export const updateToolById = async (token: string, id: string, tool: object) => return res; }; +export const updateToolAccessGrants = async ( + token: string, + id: string, + accessGrants: any[] +) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/tools/id/${id}/access/update`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify({ access_grants: accessGrants }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const deleteToolById = async (token: string, id: string) => { let error = null; diff --git a/src/lib/components/workspace/Models/ModelEditor.svelte b/src/lib/components/workspace/Models/ModelEditor.svelte index 6fd7d1ce2..a47d738f7 100644 --- a/src/lib/components/workspace/Models/ModelEditor.svelte +++ b/src/lib/components/workspace/Models/ModelEditor.svelte @@ -25,6 +25,7 @@ import PromptSuggestions from './PromptSuggestions.svelte'; import AccessControlModal from '../common/AccessControlModal.svelte'; import LockClosed from '$lib/components/icons/LockClosed.svelte'; + import { updateModelAccessGrants } from '$lib/apis/models'; const i18n = getContext('i18n'); @@ -331,6 +332,16 @@ accessRoles={preset ? ['read', 'write'] : ['read']} share={$user?.permissions?.sharing?.models || $user?.role === 'admin'} sharePublic={$user?.permissions?.sharing?.public_models || $user?.role === 'admin' || edit} + onChange={async () => { + if (edit && model?.id) { + try { + await updateModelAccessGrants(localStorage.token, model.id, accessGrants); + toast.success($i18n.t('Saved')); + } catch (error) { + toast.error(`${error}`); + } + } + }} /> {#if onBack} diff --git a/src/lib/components/workspace/Tools/ToolkitEditor.svelte b/src/lib/components/workspace/Tools/ToolkitEditor.svelte index 70533099c..4822ebbf8 100644 --- a/src/lib/components/workspace/Tools/ToolkitEditor.svelte +++ b/src/lib/components/workspace/Tools/ToolkitEditor.svelte @@ -1,10 +1,12 @@