mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-04 14:50:30 +08:00
fix: synchronize chat model configuration (#17717)
This commit is contained in:
@@ -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", []))
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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: {}
|
||||
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user