From 8ac047d11ab4cca672dec696b7fb2c4c2686a6e8 Mon Sep 17 00:00:00 2001 From: buua436 Date: Mon, 3 Aug 2026 16:43:53 +0800 Subject: [PATCH] fix: synchronize chat model configuration (#17717) --- api/apps/restful_apis/chat_api.py | 164 ++++++++++-------- api/db/services/dialog_service.py | 9 +- test/testcases/restful_api/test_chats.py | 17 +- .../test_user_tenant_routes_unit.py | 9 + .../chat/app-settings/chat-settings.tsx | 4 + .../pages/next-chats/hooks/use-create-chat.ts | 1 + 6 files changed, 126 insertions(+), 78 deletions(-) diff --git a/api/apps/restful_apis/chat_api.py b/api/apps/restful_apis/chat_api.py index 4eadcaf546..4472b4065d 100644 --- a/api/apps/restful_apis/chat_api.py +++ b/api/apps/restful_apis/chat_api.py @@ -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", [])) diff --git a/api/db/services/dialog_service.py b/api/db/services/dialog_service.py index 3a90dc29c4..aa4c3caf50 100644 --- a/api/db/services/dialog_service.py +++ b/api/db/services/dialog_service.py @@ -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 content, interleaved with the real token stream. diff --git a/test/testcases/restful_api/test_chats.py b/test/testcases/restful_api/test_chats.py index c4fad4d8da..bc75a96331 100644 --- a/test/testcases/restful_api/test_chats.py +++ b/test/testcases/restful_api/test_chats.py @@ -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) diff --git a/test/testcases/restful_api/test_user_tenant_routes_unit.py b/test/testcases/restful_api/test_user_tenant_routes_unit.py index 9047d781ac..f95b84ea71 100644 --- a/test/testcases/restful_api/test_user_tenant_routes_unit.py +++ b/test/testcases/restful_api/test_user_tenant_routes_unit.py @@ -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: {} diff --git a/web/src/pages/next-chats/chat/app-settings/chat-settings.tsx b/web/src/pages/next-chats/chat/app-settings/chat-settings.tsx index db09351335..464d935ac8 100644 --- a/web/src/pages/next-chats/chat/app-settings/chat-settings.tsx +++ b/web/src/pages/next-chats/chat/app-settings/chat-settings.tsx @@ -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', diff --git a/web/src/pages/next-chats/hooks/use-create-chat.ts b/web/src/pages/next-chats/hooks/use-create-chat.ts index 3492c6bc08..f195ff48aa 100644 --- a/web/src/pages/next-chats/hooks/use-create-chat.ts +++ b/web/src/pages/next-chats/hooks/use-create-chat.ts @@ -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,