This commit is contained in:
Timothy Jaeryang Baek 2026-07-26 23:09:22 -04:00
parent 95d590b360
commit ed663f16ec
4 changed files with 22 additions and 28 deletions

View file

@ -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)):
"""

View file

@ -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',

View file

@ -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'):

View file

@ -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