fix: synchronize chat model configuration (#17717)

This commit is contained in:
buua436
2026-08-03 16:43:53 +08:00
committed by GitHub
parent a12dee5524
commit 8ac047d11a
6 changed files with 126 additions and 78 deletions

View File

@@ -28,7 +28,14 @@ from werkzeug.exceptions import BadRequest
from api.apps import current_user, login_required
from api.apps.restful_apis._generation_params import merge_generation_config, pop_generation_config
from api.db.joint_services.tenant_model_service import get_api_key, get_tenant_default_model_by_type, resolve_model_config
from api.db.joint_services.tenant_model_service import (
get_api_key,
get_composite_model_name_by_id,
get_model_config_by_id,
get_tenant_default_model_by_type,
resolve_model_config,
resolve_model_id,
)
from api.db.services.chunk_feedback_service import ChunkFeedbackService
from api.db.services.conversation_service import ConversationService, structure_answer
from api.db.services.dialog_service import DialogService, gen_mindmap, rag_agent
@@ -268,51 +275,70 @@ def _normalize_completion_messages(req):
return (messages, msg), None
async def _validate_llm_id(llm_id, tenant_id, llm_setting=None):
if not llm_id:
def _llm_model_type(llm_setting):
model_type = (llm_setting or {}).get("model_type")
if isinstance(model_type, list):
return "vision" if "vision" in model_type else "chat"
return model_type if isinstance(model_type, str) and model_type in {"chat", "vision"} else "chat"
async def _normalize_model_pair(req, tenant_id, name_field, id_field, model_type):
"""Validate and synchronize a model name with its tenant-model ID."""
if name_field not in req and id_field not in req:
return None
conf_model_type = (llm_setting or {}).get("model_type")
if isinstance(conf_model_type, str):
model_type = conf_model_type if conf_model_type in {"chat", "vision"} else "chat"
elif isinstance(conf_model_type, list):
model_type = "vision" if "vision" in conf_model_type else "chat"
model_name = req.get(name_field)
model_id = req.get(id_field)
resolved_name_id = None
if model_name is not None and not isinstance(model_name, str):
return f"`{name_field}` must be a string"
if model_id is not None and not isinstance(model_id, str):
return f"`{id_field}` must be a tenant model id"
if model_id:
try:
await thread_pool_exec(get_model_config_by_id, tenant_id, model_type, model_id)
except LookupError as e:
logging.error("Fail to get %s config by tenant model id %s: %s", model_type, model_id, e)
return f"`{id_field}` must be a valid tenant model id"
if model_name:
if model_type == "rerank" and model_name.split("@")[0] in _DEFAULT_RERANK_MODELS:
if model_id:
return f"`{name_field}` and `{id_field}` must refer to the same model"
else:
try:
# A name field may already contain a tenant model ID. Resolve it
# strictly first, then fall back to the composite model name.
try:
await thread_pool_exec(get_model_config_by_id, tenant_id, model_type, model_name)
resolved_name_id = model_name
except LookupError:
resolved_name_id = await thread_pool_exec(resolve_model_id, tenant_id, model_type, model_name)
except LookupError as e:
logging.error("Fail to resolve %s %s: %s", name_field, model_name, e)
return f"`{name_field}` {model_name} doesn't exist"
if model_id and model_name:
if resolved_name_id != model_id:
return f"`{name_field}` and `{id_field}` must refer to the same model"
elif model_name:
req[id_field] = resolved_name_id
elif model_id:
try:
req[name_field] = await thread_pool_exec(get_composite_model_name_by_id, model_id)
except LookupError as e:
logging.error("Fail to get composite name for %s %s: %s", id_field, model_id, e)
return f"`{id_field}` must be a valid tenant model id"
else:
model_type = "chat"
try:
await thread_pool_exec(
resolve_model_config,
tenant_id=tenant_id,
model_type=model_type,
model_ref=llm_id,
)
except Exception as e:
logging.error(f"Fail to get model config for {llm_id}: {e}")
return f"`llm_id` {llm_id} doesn't exist"
# Clearing one side must not leave a stale ID/name on the other side.
req[id_field] = None
req[name_field] = ""
return None
async def _validate_rerank_id(rerank_id, tenant_id):
if not rerank_id:
return None
parts = rerank_id.split("@")
llm_name = parts[0]
if llm_name in _DEFAULT_RERANK_MODELS:
return None
try:
await thread_pool_exec(
resolve_model_config,
tenant_id=tenant_id,
model_ref=rerank_id,
model_type="rerank",
)
except Exception as e:
logging.error(f"Fail to get model config for {rerank_id}: {e}")
return f"`rerank_id` {rerank_id} doesn't exist"
return None
# def _validate_prompt_config(prompt_config):
# for parameter in prompt_config.get("parameters", []):
# if parameter.get("optional"):
@@ -387,15 +413,17 @@ async def create():
req["kb_ids"] = kb_ids
req.pop("dataset_ids", None)
if "llm_id" in req:
err = await _validate_llm_id(req.get("llm_id"), current_user.id, req.get("llm_setting"))
if err:
return get_data_error_result(message=err)
if req.get("llm_id") is None and req.get("tenant_llm_id") is None:
req["llm_id"] = tenant.tenant_llm_id
if "rerank_id" not in req and "tenant_rerank_id" not in req:
req["rerank_id"] = ""
if "rerank_id" in req:
err = await _validate_rerank_id(req.get("rerank_id"), current_user.id)
if err:
return get_data_error_result(message=err)
err = await _normalize_model_pair(req, current_user.id, "llm_id", "tenant_llm_id", _llm_model_type(req.get("llm_setting")))
if err:
return get_data_error_result(message=err)
err = await _normalize_model_pair(req, current_user.id, "rerank_id", "tenant_rerank_id", "rerank")
if err:
return get_data_error_result(message=err)
if "prompt_config" in req:
if not isinstance(req["prompt_config"], dict):
@@ -563,15 +591,13 @@ async def update_chat(chat_id):
req["kb_ids"] = kb_ids
req.pop("dataset_ids", None)
if "llm_id" in req:
err = await _validate_llm_id(req.get("llm_id"), current_user.id, req.get("llm_setting"))
if err:
return get_data_error_result(message=err)
if "rerank_id" in req:
err = await _validate_rerank_id(req.get("rerank_id"), current_user.id)
if err:
return get_data_error_result(message=err)
effective_llm_setting = req.get("llm_setting", current_chat.get("llm_setting", {}))
err = await _normalize_model_pair(req, current_user.id, "llm_id", "tenant_llm_id", _llm_model_type(effective_llm_setting))
if err:
return get_data_error_result(message=err)
err = await _normalize_model_pair(req, current_user.id, "rerank_id", "tenant_rerank_id", "rerank")
if err:
return get_data_error_result(message=err)
if "prompt_config" in req:
if not isinstance(req["prompt_config"], dict):
@@ -643,15 +669,20 @@ async def patch_chat(chat_id):
req["kb_ids"] = kb_ids
req.pop("dataset_ids", None)
if "llm_id" in req:
err = await _validate_llm_id(req.get("llm_id"), current_user.id, req.get("llm_setting"))
if err:
return get_data_error_result(message=err)
if "llm_setting" in req:
if not isinstance(req["llm_setting"], dict):
return get_data_error_result(message="`llm_setting` should be an object.")
llm_setting = deepcopy(current_chat.get("llm_setting") or {})
llm_setting.update(req["llm_setting"])
req["llm_setting"] = llm_setting
if "rerank_id" in req:
err = await _validate_rerank_id(req.get("rerank_id"), current_user.id)
if err:
return get_data_error_result(message=err)
effective_llm_setting = req.get("llm_setting", current_chat.get("llm_setting", {}))
err = await _normalize_model_pair(req, current_user.id, "llm_id", "tenant_llm_id", _llm_model_type(effective_llm_setting))
if err:
return get_data_error_result(message=err)
err = await _normalize_model_pair(req, current_user.id, "rerank_id", "tenant_rerank_id", "rerank")
if err:
return get_data_error_result(message=err)
if "prompt_config" in req:
if not isinstance(req["prompt_config"], dict):
@@ -663,11 +694,6 @@ async def patch_chat(chat_id):
# if err:
# return get_data_error_result(message=err)
if "llm_setting" in req:
llm_setting = deepcopy(current_chat.get("llm_setting", {}))
llm_setting.update(req["llm_setting"])
req["llm_setting"] = llm_setting
# if "prompt_config" in req or "kb_ids" in req:
# prompt_config = req.get("prompt_config", current_chat.get("prompt_config", {}))
# kb_ids = req.get("kb_ids", current_chat.get("kb_ids", []))

View File

@@ -1859,16 +1859,15 @@ async def gen_mindmap(question, kb_ids, tenant_id, search_config={}):
async def rag_agent(dialog, messages, stream=True, **kwargs):
logging.debug("Begin rag_agent")
prompt_config = dialog.prompt_config or {}
assert messages[-1]["role"] == "user", "The last content of this conversation is not from user."
prompt_config = dialog.prompt_config
if not prompt_config.get("reasoning", 0) and not kwargs.get("reasoning"):
async for ans in async_chat(dialog, messages, stream, **kwargs):
yield ans
return
kbs, embd_mdl, rerank_mdl, chat_mdl, tts_mdl = get_models(dialog)
use_web_search = _should_use_web_search(prompt_config, kwargs.get("internet"))
logging.debug("web_search kb=%s tavily=%s internet=%r enabled=%s", bool(dialog.kb_ids), bool(dialog.prompt_config.get("tavily_api_key")), kwargs.get("internet"), use_web_search)
logging.debug("web_search kb=%s tavily=%s internet=%r enabled=%s", bool(dialog.kb_ids), bool(prompt_config.get("tavily_api_key")), kwargs.get("internet"), use_web_search)
tenant_ids = list(set([kb.tenant_id for kb in kbs]))
# "reasoning" arrives as "1".."4" mapping to the ordered THINKING_MODES
# (low, medium, high, ultra); fall back to "medium" on anything else.
@@ -1881,12 +1880,13 @@ async def rag_agent(dialog, messages, stream=True, **kwargs):
except (TypeError, ValueError):
thinking_mode = "medium"
gen_conf = dialog.llm_setting or {}
rag_tools = RAGTools(
tenant_ids,
chat_mdl,
embed_mdl=embd_mdl,
kb_ids=dialog.kb_ids,
tav=Tavily(prompt_config["tavily_api_key"]) if use_web_search else None,
tav=Tavily(prompt_config.get("tavily_api_key")) if use_web_search else None,
do_refer=False,
thinking_mode=thinking_mode,
)
@@ -1949,7 +1949,6 @@ async def rag_agent(dialog, messages, stream=True, **kwargs):
# small models mangle or drop, so the client receives nothing.
if getattr(chat_mdl, "mdl", None) is not None:
chat_mdl.mdl.terminal_tools = {"rag"}
gen_conf = dialog.llm_setting
if stream:
# Surface the agentic pipeline's bracket-tagged progress logs to the
# client as <think> content, interleaved with the real token stream.

View File

@@ -765,6 +765,15 @@ def _load_chat_routes_unit_module(monkeypatch):
tenant_model_service_mod = ModuleType("api.db.joint_services.tenant_model_service")
tenant_model_service_mod.get_model_config_from_provider_instance = lambda *_args, **_kwargs: {}
def _get_model_config_by_id(_tenant_id, _model_type, model_ref):
if model_ref == "tenant-llm-id":
return {}
raise LookupError(f"unknown tenant model id: {model_ref}")
tenant_model_service_mod.get_model_config_by_id = _get_model_config_by_id
tenant_model_service_mod.resolve_model_id = lambda _tenant_id, _model_type, model_name: model_name
tenant_model_service_mod.get_composite_model_name_by_id = lambda model_id: model_id
tenant_model_service_mod.resolve_model_config = lambda *_args, **_kwargs: {}
tenant_model_service_mod.get_tenant_default_model_by_type = lambda *_args, **_kwargs: {}
tenant_model_service_mod.get_api_key = lambda *_args, **_kwargs: SimpleNamespace(id=1)
@@ -1149,11 +1158,11 @@ def test_chat_create_accepts_provider_scoped_rerank_id_unit(monkeypatch):
monkeypatch.setattr(module.KnowledgebaseService, "query", lambda **_kwargs: [_DummyKB()])
monkeypatch.setattr(module.KnowledgebaseService, "get_by_id", lambda _id: (True, _DummyKB()))
def _get_model_config_from_provider_instance(**kwargs):
query_calls.append(kwargs)
return {}
def _resolve_model_id(tenant_id, model_type, model_name):
query_calls.append({"tenant_id": tenant_id, "model_ref": model_name, "model_type": model_type})
return model_name
monkeypatch.setattr(module, "resolve_model_config", _get_model_config_from_provider_instance)
monkeypatch.setattr(module, "resolve_model_id", _resolve_model_id)
def _save(**kwargs):
saved.update(kwargs)

View File

@@ -1503,6 +1503,15 @@ def _load_chat_routes_unit_module(monkeypatch):
tenant_model_provider_mod = ModuleType("api.db.joint_services.tenant_model_service")
tenant_model_provider_mod.get_model_config_from_provider_instance = lambda *_args, **_kwargs: {}
def _get_model_config_by_id(_tenant_id, _model_type, model_ref):
if model_ref == "tenant-llm-id":
return {}
raise LookupError(f"unknown tenant model id: {model_ref}")
tenant_model_provider_mod.get_model_config_by_id = _get_model_config_by_id
tenant_model_provider_mod.resolve_model_id = lambda _tenant_id, _model_type, model_name: model_name
tenant_model_provider_mod.get_composite_model_name_by_id = lambda model_id: model_id
tenant_model_provider_mod.resolve_model_config = lambda *_args, **_kwargs: {}
tenant_model_provider_mod.get_tenant_default_model_by_type = lambda *_args, **_kwargs: {}

View File

@@ -87,6 +87,10 @@ export function ChatSettings({ hasSingleChatBox }: ChatSettingsProps) {
// Add model_type to llm_setting based on the selected llm_id
if (nextValues.llm_id) {
// The model selector returns the tenant model ID. Keep the legacy
// llm_id and the tenant-scoped ID synchronized; the backend gives
// tenant_llm_id precedence when resolving the chat model.
nextValues.tenant_llm_id = nextValues.llm_id;
nextValues.llm_setting = {
...nextValues.llm_setting,
model_type: findLlmByUuid(nextValues.llm_id)?.model_type || 'chat',

View File

@@ -39,6 +39,7 @@ export const useCreateChatDialog = () => {
toc_enhance: false,
},
llm_id: defaultModelDictionary?.llm_id,
tenant_llm_id: defaultModelDictionary?.llm_id,
llm_setting: {},
similarity_threshold: 0.2,
vector_similarity_weight: 0.3,