mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-13 01:02:25 -06:00
refac
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
@@ -240,6 +241,7 @@ from open_webui.utils.middleware import (
|
||||
process_chat_payload,
|
||||
process_chat_response,
|
||||
)
|
||||
from open_webui.utils.misc import merge_model_params
|
||||
from open_webui.utils.model_ids import strip_provider_model_prefix
|
||||
from open_webui.utils.models import (
|
||||
check_model_access,
|
||||
@@ -1097,17 +1099,14 @@ async def chat_completion(
|
||||
await _set_direct_model(request, model, user)
|
||||
|
||||
# Model params: global defaults as base, per-model overrides win
|
||||
default_model_params = await Config.get('models.default_params', {}) or {}
|
||||
model_info_params = {
|
||||
**default_model_params,
|
||||
**(model_info.params.model_dump() if model_info and model_info.params else {}),
|
||||
}
|
||||
default_model_params = copy.deepcopy(await Config.get('models.default_params', {}) or {})
|
||||
model_info_params = merge_model_params(
|
||||
default_model_params,
|
||||
model_info.params.model_dump() if model_info and model_info.params else {},
|
||||
)
|
||||
request_params = {key: value for key, value in (form_data.get('params') or {}).items() if value is not None}
|
||||
if model_info_params or request_params:
|
||||
form_data['params'] = {
|
||||
**model_info_params,
|
||||
**request_params,
|
||||
}
|
||||
form_data['params'] = merge_model_params(model_info_params, request_params)
|
||||
|
||||
# Check base model existence for custom models
|
||||
if model_info and model_info.base_model_id:
|
||||
|
||||
@@ -1976,14 +1976,14 @@ def apply_params_to_form_data(form_data, model):
|
||||
|
||||
if model.get('owned_by') == 'ollama':
|
||||
# Ollama specific parameters
|
||||
form_data['options'] = params
|
||||
form_data['options'] = {**params, **(form_data.get('options') or {})}
|
||||
else:
|
||||
if isinstance(params, dict):
|
||||
for key, value in params.items():
|
||||
if value is not None:
|
||||
if value is not None and key not in form_data:
|
||||
form_data[key] = value
|
||||
|
||||
if 'logit_bias' in params and params['logit_bias'] is not None:
|
||||
if 'logit_bias' in params and params['logit_bias'] is not None and 'logit_bias' not in form_data:
|
||||
try:
|
||||
logit_bias = convert_logit_bias_input_to_json(params['logit_bias'])
|
||||
|
||||
|
||||
@@ -29,6 +29,15 @@ def deep_update(d, u):
|
||||
return d
|
||||
|
||||
|
||||
def merge_model_params(base: dict, override: dict) -> dict:
|
||||
params = {**base, **override}
|
||||
base_custom = base.get('custom_params')
|
||||
override_custom = override.get('custom_params')
|
||||
if isinstance(base_custom, dict) and (override_custom is None or isinstance(override_custom, dict)):
|
||||
params['custom_params'] = {**base_custom, **(override_custom or {})}
|
||||
return params
|
||||
|
||||
|
||||
def _strip_filter_entry(entry):
|
||||
# Compose list-form env syntax passes surrounding quotes through verbatim
|
||||
return (entry or '').strip().strip('"\'').strip()
|
||||
|
||||
@@ -67,7 +67,7 @@ def apply_model_params_to_body(params: dict, form_data: dict, mappings: dict[str
|
||||
return form_data
|
||||
|
||||
for key, value in params.items():
|
||||
if value is not None:
|
||||
if value is not None and key not in form_data:
|
||||
if key in mappings:
|
||||
cast_func = mappings[key]
|
||||
if isinstance(cast_func, Callable):
|
||||
|
||||
Reference in New Issue
Block a user