diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 409bbb6640..6b9269e8fe 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -236,6 +236,7 @@ from open_webui.utils.middleware import ( process_chat_payload, process_chat_response, ) +from open_webui.utils.model_ids import strip_provider_model_prefix from open_webui.utils.models import ( check_model_access, get_all_base_models, @@ -888,12 +889,6 @@ class ModelUnloadForm(BaseModel): model: str -def strip_provider_model_prefix(model_id: str, prefix_id: str | None) -> str: - if prefix_id and model_id.startswith(f'{prefix_id}.'): - return model_id[len(f'{prefix_id}.') :] - return model_id - - @app.post('/api/models/unload') async def unload_model(request: Request, form_data: ModelUnloadForm, user=Depends(get_admin_user)): """ diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index ff1bc3a19a..09e68d540f 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -41,6 +41,7 @@ from open_webui.models.users import UserModel from open_webui.utils.access_control import check_model_access from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.headers import get_custom_headers, include_user_info_headers +from open_webui.utils.model_ids import strip_provider_model_prefix from open_webui.utils.misc import calculate_sha256 from open_webui.utils.payload import ( apply_model_params_to_body_ollama, @@ -895,8 +896,7 @@ async def embed( key = get_api_key(url_idx, url, api_configs) prefix_id = api_config.get('prefix_id') - if prefix_id: - form_data.model = form_data.model.replace(f'{prefix_id}.', '') + form_data.model = strip_provider_model_prefix(form_data.model, prefix_id) return await send_request( f'{url}/api/embed', @@ -947,8 +947,7 @@ async def embeddings( key = get_api_key(url_idx, url, api_configs) prefix_id = api_config.get('prefix_id') - if prefix_id: - form_data.model = form_data.model.replace(f'{prefix_id}.', '') + form_data.model = strip_provider_model_prefix(form_data.model, prefix_id) return await send_request( f'{url}/api/embeddings', @@ -1003,8 +1002,7 @@ async def generate_completion( api_config = api_configs.get(str(url_idx), api_configs.get(url, {})) prefix_id = api_config.get('prefix_id') - if prefix_id: - form_data.model = form_data.model.replace(f'{prefix_id}.', '') + form_data.model = strip_provider_model_prefix(form_data.model, prefix_id) return await send_request( f'{url}/api/generate', @@ -1063,6 +1061,10 @@ async def get_ollama_url(request: Request, model: str, url_idx: int | None = Non await validate_ollama_backend_idx(request, model, url_idx, user) if url_idx is None: models = request.app.state.OLLAMA_MODELS + if not models or model not in models: + await get_all_models.cache.clear() + await get_all_models(request, user=user) + models = request.app.state.OLLAMA_MODELS if model not in models: raise HTTPException( status_code=400, @@ -1135,8 +1137,7 @@ async def generate_chat_completion( api_config = resolve_api_config(api_configs, url_idx, url) prefix_id = api_config.get('prefix_id') - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) return await send_request( f'{url}/api/chat', @@ -1232,8 +1233,7 @@ async def generate_openai_completion( api_config = resolve_api_config(api_configs, url_idx, url) prefix_id = api_config.get('prefix_id') - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) return await send_request( f'{url}/v1/completions', @@ -1283,8 +1283,7 @@ async def generate_openai_embeddings( api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url) prefix_id = api_config.get('prefix_id') - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) return await send_request( f'{url}/v1/embeddings', @@ -1342,8 +1341,7 @@ async def generate_openai_chat_completion( api_config = resolve_api_config(api_configs, url_idx, url) prefix_id = api_config.get('prefix_id') - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) return await send_request( f'{url}/v1/chat/completions', @@ -1394,8 +1392,7 @@ async def generate_anthropic_messages( api_config = api_configs.get(str(url_idx), api_configs.get(url, {})) # Legacy support prefix_id = api_config.get('prefix_id', None) - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) return await send_request( f'{url}/v1/messages', @@ -1452,8 +1449,7 @@ async def generate_responses( api_config = api_configs.get(str(url_idx), api_configs.get(url, {})) # Legacy support prefix_id = api_config.get('prefix_id', None) - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) return await send_request( f'{url}/v1/responses', diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 5d0f9062ed..d331f9f690 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -44,6 +44,7 @@ from open_webui.utils.access_control import check_model_access, has_connection_a from open_webui.utils.anthropic import get_anthropic_models, is_anthropic_url from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.headers import get_custom_headers, include_user_info_headers +from open_webui.utils.model_ids import strip_provider_model_prefix from open_webui.utils.misc import ( convert_logit_bias_input_to_json, stream_chunks_handler, @@ -328,8 +329,7 @@ async def get_anthropic_token_count_target(request: Request, form_data: dict, us url, key, api_config = await get_openai_connection(model['urlIdx']) prefix_id = api_config.get('prefix_id') - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user) return requested_model, payload, url, key, headers, cookies @@ -1268,8 +1268,7 @@ async def generate_chat_completion( url, key, api_config = await get_openai_connection(idx) prefix_id = api_config.get('prefix_id', None) - if prefix_id: - payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + payload['model'] = strip_provider_model_prefix(payload['model'], prefix_id) # Add user info to the payload if the model is a pipeline if 'pipeline' in model and model.get('pipeline'): diff --git a/backend/open_webui/utils/model_ids.py b/backend/open_webui/utils/model_ids.py new file mode 100644 index 0000000000..c28c8b532d --- /dev/null +++ b/backend/open_webui/utils/model_ids.py @@ -0,0 +1,4 @@ +def strip_provider_model_prefix(model_id: str, prefix_id: str | None) -> str: + if prefix_id and model_id.startswith(f'{prefix_id}.'): + return model_id[len(f'{prefix_id}.') :] + return model_id