diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index e5c9f2e7e8..619f8e43a0 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -246,9 +246,48 @@ router = APIRouter() LLAMACPP_LOADED_STATES = {'loaded', 'sleeping'} LLAMACPP_UNLOADED_STATES = {'loading', 'unloaded'} +MODEL_MANAGEMENT_ENDPOINTS = { + 'llama.cpp': { + 'list': '/models', + 'download': '/models', + 'delete': '/models', + 'load': '/models/load', + 'unload': '/models/unload', + 'sse': '/models/sse', + }, + 'lmstudio': { + 'list': '/api/v1/models', + 'download': '/api/v1/models/download', + 'download_status': '/api/v1/models/download/status/{job_id}', + 'load': '/api/v1/models/load', + 'unload': '/api/v1/models/unload', + }, +} -def get_llamacpp_model_loaded_state(model: dict, provider: str, manual_model_ids: bool = False) -> bool | None: +def get_model_management_root_url(url: str, provider: str) -> str: + root_url = url.rstrip('/') + if provider in ('llama.cpp', 'lmstudio'): + for suffix in ('/api/v1', '/api/v0', '/v1'): + if root_url.endswith(suffix): + return root_url.removesuffix(suffix) + + return root_url + + +def get_provider_model_loaded_state(model: dict, provider: str, manual_model_ids: bool = False) -> bool | None: + if provider == 'lmstudio': + if model.get('loaded_instances'): + return True + + state = model.get('state') + if state == 'loaded': + return True + if state == 'not-loaded': + return False + + return None + if provider != 'llama.cpp': return None @@ -307,6 +346,111 @@ async def get_openai_connection(idx: int) -> tuple[str, str, dict]: return url, key, api_config +async def clear_openai_model_cache(request: Request): + await get_all_models.cache.clear() + request.app.state.BASE_MODELS = [] + request.app.state.OPENAI_MODELS = {} + models = getattr(request.app.state, 'MODELS', None) + if hasattr(models, 'clear'): + models.clear() + else: + request.app.state.MODELS = {} + + +async def get_model_management_connection(url_idx: int) -> tuple[str, str, dict, str]: + if not await Config.get('openai.enable'): + raise HTTPException(status_code=503, detail='OpenAI API is disabled') + + try: + url, key, api_config = await get_openai_connection(url_idx) + except IndexError: + raise HTTPException(status_code=404, detail='Connection not found') + + provider = api_config.get('provider', '') + if provider not in MODEL_MANAGEMENT_ENDPOINTS: + raise HTTPException( + status_code=400, + detail=f'Provider "{provider or "default"}" does not support model management', + ) + + return get_model_management_root_url(url, provider), key, api_config, provider + + +def get_model_management_path(provider: str, operation: str, path_params: dict | None = None) -> str: + try: + path = MODEL_MANAGEMENT_ENDPOINTS[provider][operation] + except KeyError: + raise HTTPException(status_code=400, detail=f'Provider "{provider}" does not support {operation}') + + return path.format(**(path_params or {})) + + +def get_model_management_payload(provider: str, operation: str, payload: dict | None) -> dict | None: + if provider == 'lmstudio' and operation == 'unload' and payload: + return {'instance_id': payload.get('instance_id') or payload.get('model')} + + return payload + + +async def send_model_management_request( + request: Request, + url_idx: int, + operation: str, + method: str = 'GET', + payload: dict | None = None, + query: dict | None = None, + path_params: dict | None = None, + stream: bool = False, + user: UserModel | None = None, +): + root_url, key, api_config, provider = await get_model_management_connection(url_idx) + path = get_model_management_path(provider, operation, path_params=path_params) + payload = get_model_management_payload(provider, operation, payload) + headers, cookies = await get_headers_and_cookies(request, root_url, key, api_config, user=user) + + response = None + streaming = False + try: + session = await get_session() + response = await session.request( + method, + f'{root_url}{path}', + json=payload, + params=query, + headers=headers, + cookies=cookies, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + timeout=get_client_timeout(stream=stream), + ) + + if not response.ok: + try: + error = await response.json(loads=JSONCodec.loads) + except Exception: + error = await response.text() + raise HTTPException(status_code=response.status, detail=error) + + if stream: + streaming = True + return StreamingResponse( + stream_wrapper(response, passthrough=True), + status_code=response.status, + headers=_clean_proxy_headers(response.headers), + ) + + try: + return await response.json(loads=JSONCodec.loads) + except Exception: + return {'success': True} + except HTTPException: + raise + except Exception as e: + raise HTTPException(status_code=response.status if response else 500, detail=str(e)) + finally: + if not streaming: + await cleanup_response(response) + + async def get_anthropic_token_count_target(request: Request, form_data: dict, user: UserModel): """Resolve the upstream LiteLLM connection for an Anthropic token-count request.""" requested_model = form_data.get('model') @@ -697,7 +841,7 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]: 'urlIdx': idx, } - loaded = get_llamacpp_model_loaded_state( + loaded = get_provider_model_loaded_state( model, provider, manual_model_ids=bool(api_config.get('model_ids')), @@ -793,6 +937,121 @@ async def get_models(request: Request, url_idx: int | None = None, user=Depends( return models +class ProviderModelOperationForm(BaseModel): + model: str + model_config = ConfigDict(extra='allow') + + +@router.get('/models/{url_idx}/catalog') +async def get_provider_model_catalog(request: Request, url_idx: int, user=Depends(get_admin_user)): + return await send_model_management_request(request, url_idx, 'list', user=user) + + +@router.post('/models/{url_idx}/download') +async def download_provider_model( + request: Request, + url_idx: int, + form_data: ProviderModelOperationForm, + user=Depends(get_admin_user), +): + root_url, _, api_config, provider = await get_model_management_connection(url_idx) + payload = form_data.model_dump(exclude_none=True) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + + result = await send_model_management_request(request, url_idx, 'download', 'POST', payload, user=user) + await clear_openai_model_cache(request) + await publish_event( + request, + EVENTS.MODEL_PROVIDER_MODEL_CREATED, + actor=user, + subject_id=payload['model'], + data={'provider': provider, 'url_idx': url_idx, 'base_url': root_url}, + ) + return result + + +@router.get('/models/{url_idx}/download/status/{job_id}') +async def get_provider_model_download_status( + request: Request, + url_idx: int, + job_id: str, + user=Depends(get_admin_user), +): + return await send_model_management_request( + request, + url_idx, + 'download_status', + path_params={'job_id': job_id}, + user=user, + ) + + +@router.post('/models/{url_idx}/load') +async def load_provider_model( + request: Request, + url_idx: int, + form_data: ProviderModelOperationForm, + user=Depends(get_admin_user), +): + _, _, api_config, _ = await get_model_management_connection(url_idx) + payload = form_data.model_dump(exclude_none=True) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + + result = await send_model_management_request(request, url_idx, 'load', 'POST', payload, user=user) + await clear_openai_model_cache(request) + return result + + +@router.post('/models/{url_idx}/unload') +async def unload_provider_model( + request: Request, + url_idx: int, + form_data: ProviderModelOperationForm, + user=Depends(get_admin_user), +): + _, _, api_config, _ = await get_model_management_connection(url_idx) + payload = form_data.model_dump(exclude_none=True) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + + result = await send_model_management_request(request, url_idx, 'unload', 'POST', payload, user=user) + await clear_openai_model_cache(request) + return result + + +@router.get('/models/{url_idx}/sse') +async def stream_provider_model_events(request: Request, url_idx: int, user=Depends(get_admin_user)): + return await send_model_management_request(request, url_idx, 'sse', stream=True, user=user) + + +@router.delete('/models/{url_idx}') +async def delete_provider_model( + request: Request, + url_idx: int, + model: str, + user=Depends(get_admin_user), +): + root_url, _, api_config, provider = await get_model_management_connection(url_idx) + actual_model = strip_provider_model_prefix(model, api_config.get('prefix_id')) + + result = await send_model_management_request( + request, + url_idx, + 'delete', + 'DELETE', + query={'model': actual_model}, + user=user, + ) + await clear_openai_model_cache(request) + await publish_event( + request, + EVENTS.MODEL_PROVIDER_MODEL_DELETED, + actor=user, + subject_id=actual_model, + data={'provider': provider, 'url_idx': url_idx, 'base_url': root_url}, + ) + return result + + class ConnectionVerificationForm(BaseModel): url: str key: str diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index 4cdefd0de5..1c648b7f38 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -71,6 +71,10 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) 'evaluation.arena.models', 'models.default_metadata', ) + if refresh: + await openai.get_all_models.cache.clear() + await ollama.get_all_models.cache.clear() + if ( request.app.state.MODELS and request.app.state.BASE_MODELS diff --git a/src/lib/apis/openai/index.ts b/src/lib/apis/openai/index.ts index d18565fec3..28bd4c299a 100644 --- a/src/lib/apis/openai/index.ts +++ b/src/lib/apis/openai/index.ts @@ -131,9 +131,185 @@ export const getOpenAIModels = async (token: string, urlIdx?: number) => { return res; }; +export const getProviderModelCatalog = async (token: string, urlIdx: number) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/catalog`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + ...(token && { authorization: `Bearer ${token}` }) + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err?.detail ?? err?.error?.message ?? 'Server connection failed'; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const downloadProviderModel = async (token: string, urlIdx: number, model: string) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/download`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ model }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err?.detail ?? err?.error?.message ?? 'Server connection failed'; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getProviderModelDownloadStatus = async (token: string, urlIdx: number, jobId: string) => { + let error = null; + + const res = await fetch( + `${OPENAI_API_BASE_URL}/models/${urlIdx}/download/status/${encodeURIComponent(jobId)}`, + { + 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(); + }) + .catch((err) => { + error = err?.detail ?? err?.error?.message ?? 'Server connection failed'; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const loadProviderModel = async (token: string, urlIdx: number, model: string) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/load`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ model }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err?.detail ?? err?.error?.message ?? 'Server connection failed'; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const unloadProviderModel = async ( + token: string, + urlIdx: number, + model: string, + instanceId?: string +) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/unload`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ model, ...(instanceId ? { instance_id: instanceId } : {}) }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err?.detail ?? err?.error?.message ?? 'Server connection failed'; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const deleteProviderModel = async (token: string, urlIdx: number, model: string) => { + let error = null; + + const res = await fetch( + `${OPENAI_API_BASE_URL}/models/${urlIdx}?${new URLSearchParams({ model })}`, + { + method: 'DELETE', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err?.detail ?? err?.error?.message ?? 'Server connection failed'; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const verifyOpenAIConnection = async ( token: string = '', - connection: dict = {}, + connection: Record = {}, direct: boolean = false ) => { const { url, key, config } = connection; diff --git a/src/lib/components/AddConnectionModal.svelte b/src/lib/components/AddConnectionModal.svelte index 68eb96b70f..2c7ed22a87 100644 --- a/src/lib/components/AddConnectionModal.svelte +++ b/src/lib/components/AddConnectionModal.svelte @@ -606,6 +606,7 @@ + diff --git a/src/lib/components/admin/Settings/Connections.svelte b/src/lib/components/admin/Settings/Connections.svelte index f7b43bfb4b..fe5278d703 100644 --- a/src/lib/components/admin/Settings/Connections.svelte +++ b/src/lib/components/admin/Settings/Connections.svelte @@ -14,6 +14,7 @@ import Switch from '$lib/components/common/Switch.svelte'; import Spinner from '$lib/components/common/Spinner.svelte'; import Tooltip from '$lib/components/common/Tooltip.svelte'; + import ArrowPath from '$lib/components/icons/ArrowPath.svelte'; import Plus from '$lib/components/icons/Plus.svelte'; import OpenAIConnection from './Connections/OpenAIConnection.svelte'; @@ -50,6 +51,7 @@ let pipelineUrls: Record = {}; let showAddOpenAIConnectionModal = false; let showAddOllamaConnectionModal = false; + let modelListRefreshing = false; const updateOpenAIHandler = async () => { if (ENABLE_OPENAI_API !== null) { @@ -120,6 +122,19 @@ } }; + const refreshModelListHandler = async () => { + modelListRefreshing = true; + + try { + await models.set(await getModels()); + toast.success($i18n.t('Model list refreshed')); + } catch (error) { + toast.error(`${error}`); + } finally { + modelListRefreshing = false; + } + }; + const addOpenAIConnectionHandler = async (connection: any) => { OPENAI_API_BASE_URLS = [...OPENAI_API_BASE_URLS, connection.url]; OPENAI_API_KEYS = [...OPENAI_API_KEYS, connection.key]; @@ -373,13 +388,33 @@ )} let:labelId > - { - updateConnectionsHandler(); - }} - ariaLabelledbyId={labelId} - /> +
+ {#if connectionsConfig.ENABLE_BASE_MODELS_CACHE} + + + + {/if} + + { + updateConnectionsHandler(); + }} + ariaLabelledbyId={labelId} + /> +
{:else} diff --git a/src/lib/components/admin/Settings/Models/Manage/ManageMultipleProviderModels.svelte b/src/lib/components/admin/Settings/Models/Manage/ManageMultipleProviderModels.svelte new file mode 100644 index 0000000000..4dfdabf144 --- /dev/null +++ b/src/lib/components/admin/Settings/Models/Manage/ManageMultipleProviderModels.svelte @@ -0,0 +1,56 @@ + + +{#if connections.length > 0} +
{$i18n.t('Model providers')}
+ +
+ + {#each connections as connection} + + {/each} + +
+ +
+ +
+{/if} diff --git a/src/lib/components/admin/Settings/Models/Manage/ManageProviderModels.svelte b/src/lib/components/admin/Settings/Models/Manage/ManageProviderModels.svelte new file mode 100644 index 0000000000..26d32c5633 --- /dev/null +++ b/src/lib/components/admin/Settings/Models/Manage/ManageProviderModels.svelte @@ -0,0 +1,295 @@ + + + + +
+
+
{providerLabel || provider}
+ + + +
+ +
+ + + + +
+ + {#if loading} +
+ +
+ {:else if providerModels.length === 0} +
+ {$i18n.t('No models found')} +
+ {:else} +
+ {#each providerModels as model} + {@const modelId = getModelId(model)} + {@const displayName = getDisplayName(model)} + {@const status = getStatus(model)} +
+
+
+ {displayName} +
+ {#if displayName !== modelId} +
{modelId}
+ {/if} +
+ + {status} + +
+
+ +
+ + + + + + + + + {#if supportsDelete} + + + + {/if} +
+
+ {/each} +
+ {/if} +
diff --git a/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte b/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte index ac68484c42..1a330c2b1c 100644 --- a/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte +++ b/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte @@ -1,38 +1,72 @@ - @@ -64,8 +98,22 @@ {:else if selected !== null}
+ {#if hasOllamaManagement && hasProviderManagement} +
+ + + + +
+ {/if} {#if selected === 'ollama'} + {:else if selected === 'provider'} + {/if}
diff --git a/src/lib/i18n/locales/en-GB/translation.json b/src/lib/i18n/locales/en-GB/translation.json index 7f67b5f81f..f543d59b18 100644 --- a/src/lib/i18n/locales/en-GB/translation.json +++ b/src/lib/i18n/locales/en-GB/translation.json @@ -1716,6 +1716,7 @@ "Model Parameters": "", "Model Params": "", "Model Permissions": "", + "Model list refreshed": "", "Model removed from pinned models": "", "Model removed from selected models": "", "Model Response Mode": "", diff --git a/src/lib/i18n/locales/en-US/translation.json b/src/lib/i18n/locales/en-US/translation.json index a01be4d9f3..1e04d27cca 100644 --- a/src/lib/i18n/locales/en-US/translation.json +++ b/src/lib/i18n/locales/en-US/translation.json @@ -1719,6 +1719,7 @@ "Model Parameters": "", "Model Params": "", "Model Permissions": "", + "Model list refreshed": "", "Model removed from pinned models": "", "Model removed from selected models": "", "Model Response Mode": "",